MCPcopy Create free account
hub / github.com/alinlab/SelfPatch / train_dino

Function train_dino

main_selfpatch.py:151–314  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

149 return [self.c_head(logits[:,:1]), self.p_head(logits[:,1:])]
150
151def train_dino(args):
152 # os.environ['MASTER_PORT'] = '30000'
153 utils.init_distributed_mode(args)
154 utils.fix_random_seeds(args.seed)
155 print("git:\n {}\n".format(utils.get_sha()))
156 print("\n".join("%s: %s" % (k, str(v)) for k, v in sorted(dict(vars(args)).items())))
157 cudnn.benchmark = True
158
159 # ============ preparing data ... ============
160 transform = DataAugmentationDINO(
161 args.global_crops_scale,
162 args.local_crops_scale,
163 args.local_crops_number,
164 )
165 dataset = datasets.ImageFolder(args.data_path, transform=transform)
166 sampler = torch.utils.data.DistributedSampler(dataset, shuffle=True)
167 data_loader = torch.utils.data.DataLoader(
168 dataset,
169 sampler=sampler,
170 batch_size=args.batch_size_per_gpu,
171 num_workers=args.num_workers,
172 pin_memory=True,
173 drop_last=True,
174 )
175 print(f"Data loaded: there are {len(dataset)} images.")
176
177 # ============ building student and teacher networks ... ============
178 # we changed the name DeiT-S for ViT-S to avoid confusions
179 args.arch = args.arch.replace("deit", "vit")
180 # if the network is a vision transformer (i.e. vit_tiny, vit_small, vit_base)
181 if args.arch in vits.__dict__.keys():
182 student = vits.__dict__[args.arch](
183 patch_size=args.patch_size,
184 drop_path_rate=0.1, # stochastic depth
185 )
186 teacher = vits.__dict__[args.arch](patch_size=args.patch_size)
187 embed_dim = student.embed_dim
188 num_heads = student.num_heads
189 # otherwise, we check if the architecture is in torchvision models
190 elif args.arch in torchvision_models.__dict__.keys():
191 student = torchvision_models.__dict__[args.arch]()
192 teacher = torchvision_models.__dict__[args.arch]()
193 embed_dim = student.fc.weight.shape[1]
194 else:
195 print(f"Unknow architecture: {args.arch}")
196
197 # multi-crop wrapper handles forward with inputs of different resolutions
198 student = MultiCropWrapper(student,
199 SelfPatchHead(embed_dim,num_heads, args.k_num),
200 DINOHead(embed_dim,args.out_dim,use_bn=args.use_bn_in_head,norm_last_layer=args.norm_last_layer),
201 DINOHead(embed_dim,args.out_dim_selfpatch,use_bn=args.use_bn_in_head,norm_last_layer=args.norm_last_layer),
202 )
203 teacher = MultiCropWrapper(
204 teacher,
205 SelfPatchHead(embed_dim,num_heads, args.k_num),
206 DINOHead(embed_dim,args.out_dim, args.use_bn_in_head),
207 DINOHead(embed_dim, args.out_dim_selfpatch, args.use_bn_in_head),
208 )

Callers 1

main_selfpatch.pyFile · 0.85

Calls 7

SelfPatchHeadClass · 0.90
DINOHeadClass · 0.90
printFunction · 0.85
DINOLossClass · 0.85
train_one_epochFunction · 0.85
MultiCropWrapperClass · 0.70

Tested by

no test coverage detected