(args)
| 149 | return [self.c_head(logits[:,:1]), self.p_head(logits[:,1:])] |
| 150 | |
| 151 | def 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 | ) |
no test coverage detected