1818from ..inference .inference import get_model_path , compute_scale_from_voxel_size , get_available_models
1919from ..inference .util import _Scaler
2020
21+ # configure weak augmentations
22+ from torch_em .transform .invertible_augmentations import DEFAULT_WEAK_AUGMENTATIONS
23+
24+ # TODO - test settings - fixed kernel size for `RandomGaussianBlur`, fixed std for `RandomGaussianBlur`
25+ DEFAULT_WEAK_AUGMENTATIONS ["intensity" ] = {
26+ "RandomGaussianBlur" : {"kernel_size" : (19 , 19 ), "sigma" : (0.1 , 3.0 )},
27+ "RandomGaussianNoise" : {"mean" : (0.0 ), "std" : (0.1 )},
28+ }
2129
2230def mean_teacher_adaptation (
2331 name : str ,
@@ -40,7 +48,8 @@ def mean_teacher_adaptation(
4048 train_mask_paths : Optional [Tuple [str ]] = None ,
4149 val_mask_paths : Optional [Tuple [str ]] = None ,
4250 sample_mask_key : Optional [str ] = None ,
43- patch_sampler : Optional [callable ] = None ,
51+ unsupervised_sampler : Optional [callable ] = None ,
52+ supervised_sampler : Optional [callable ] = None ,
4453 check : bool = False ,
4554) -> None :
4655 """Run domain adaptation to transfer a network trained on a source domain for a supervised
@@ -86,10 +95,12 @@ def mean_teacher_adaptation(
8695 based on the patch_shape and size of the volumes used for training.
8796 n_samples_val: The number of val samples per epoch. By default this will be estimated
8897 based on the patch_shape and size of the volumes used for validation.
89- train_mask_paths: Sample masks used by the patch sampler to accept or reject patches for training.
90- val_mask_paths: Sample masks used by the patch sampler to accept or reject patches for validation.
98+ train_mask_paths: Sample masks used by the unsupervised sampler to accept or reject patches for training.
99+ val_mask_paths: Sample masks used by the unsupervised sampler to accept or reject patches for validation.
91100 sample_mask_key: The key to the sample mask dataset inside each file.
92- patch_sampler: Accept or reject patches based on a condition.
101+ unsupervised_sampler: Sampler to accept or reject patches for the unsupervised data stream.
102+ supervised_sampler: Sampler to accept or reject patches for the supervised data stream.
103+ Pass `False` to disable.
93104 check: Whether to check the training and validation loaders instead of running training.
94105 """ # noqa
95106 assert (supervised_train_paths is None ) == (supervised_val_paths is None )
@@ -118,8 +129,11 @@ def mean_teacher_adaptation(
118129
119130 # self training functionality
120131 pseudo_labeler = self_training .DefaultPseudoLabeler (confidence_threshold = confidence_threshold )
121- loss = self_training .DefaultSelfTrainingLoss ()
122- loss_and_metric = self_training .DefaultSelfTrainingLossAndMetric ()
132+ loss = self_training .SelfTrainingLossWithInvertibleAugmentations ()
133+ loss_and_metric = self_training .SelfTrainingLossAndMetricWithInvertibleAugmentations ()
134+
135+ ndim = 2 if is_2d else 3
136+ augmenters = torch_em .transform .invertible_augmentations .MeanTeacherAugmenters (ndim = ndim )
123137
124138 unsupervised_train_loader = get_unsupervised_loader (
125139 data_paths = unsupervised_train_paths ,
@@ -129,7 +143,7 @@ def mean_teacher_adaptation(
129143 n_samples = n_samples_train ,
130144 sample_mask_paths = train_mask_paths ,
131145 sample_mask_key = sample_mask_key ,
132- sampler = patch_sampler ,
146+ sampler = unsupervised_sampler ,
133147 )
134148 unsupervised_val_loader = get_unsupervised_loader (
135149 data_paths = unsupervised_val_paths ,
@@ -139,18 +153,20 @@ def mean_teacher_adaptation(
139153 n_samples = n_samples_val ,
140154 sample_mask_paths = val_mask_paths ,
141155 sample_mask_key = sample_mask_key ,
142- sampler = patch_sampler ,
156+ sampler = unsupervised_sampler ,
143157 )
144158
145159 if supervised_train_paths is not None :
146160 assert label_key is not None
147161 supervised_train_loader = get_supervised_loader (
148162 supervised_train_paths , raw_key_supervised , label_key ,
149163 patch_shape , batch_size , n_samples = n_samples_train ,
164+ sampler = supervised_sampler ,
150165 )
151166 supervised_val_loader = get_supervised_loader (
152167 supervised_val_paths , raw_key_supervised , label_key ,
153168 patch_shape , batch_size , n_samples = n_samples_val ,
169+ sampler = supervised_sampler ,
154170 )
155171 else :
156172 supervised_train_loader = None
@@ -166,7 +182,7 @@ def mean_teacher_adaptation(
166182 return
167183
168184 device = torch .device ("cuda" ) if torch .cuda .is_available () else torch .device ("cpu" )
169- trainer = self_training .MeanTeacherTrainer (
185+ trainer = self_training .MeanTeacherTrainerWithInvertibleAugmentations (
170186 name = name ,
171187 model = model ,
172188 optimizer = optimizer ,
@@ -187,6 +203,7 @@ def mean_teacher_adaptation(
187203 device = device ,
188204 reinit_teacher = reinit_teacher ,
189205 save_root = save_root ,
206+ augmenter = augmenters ,
190207 )
191208 trainer .fit (n_iterations )
192209
0 commit comments