(self)
| 192 | |
| 193 | |
| 194 | def _init_vision_components(self): |
| 195 | |
| 196 | vit_kwargs = self.vtp_config.vtp_model.vision_encoder |
| 197 | vit_kwargs['img_size'] = self.vtp_config.data.image_size |
| 198 | vit_kwargs['use_mask_token'] = self.vtp_config.training.train_dinov2 |
| 199 | |
| 200 | if self.vtp_config.vtp_model.vision_encoder.model_type == 'dinov3': |
| 201 | vit_kwargs['norm_layer'] = vit_kwargs.pop('norm_type') |
| 202 | vit_kwargs['ffn_ratio'] = vit_kwargs.pop('mlp_ratio') |
| 203 | vit_kwargs['layerscale_init'] = vit_kwargs.pop('init_values') |
| 204 | self.trunk = DinoVisionTransformerWithBottleneck(**vit_kwargs) |
| 205 | # Allow dropout for DINOv3 runtime forward (clip/ssl/rec) |
| 206 | self.clip_drop_rate = self.vtp_config.training.clip_drop_rate |
| 207 | self.ssl_drop_rate = self.vtp_config.training.ssl_drop_rate |
| 208 | self.rec_drop_rate = self.vtp_config.training.rec_drop_rate |
| 209 | logging.info( |
| 210 | f"DINOv3 drop configured - clip: {self.clip_drop_rate}, ssl(student): {self.ssl_drop_rate}, rec: {self.rec_drop_rate}; ssl(teacher): 0.0" |
| 211 | ) |
| 212 | else: |
| 213 | raise ValueError(f"Unsupported vision encoder: {self.vtp_config.vtp_model.vision_encoder.model_type}") |
| 214 | |
| 215 | effective_embed_dim = self.trunk.vit_feature_bottleneck |
| 216 | if self.vtp_config.training.train_clip: |
| 217 | self.proj = nn.Linear(effective_embed_dim if not self.vtp_config.vtp_model.vision_encoder.bottleneck_ae_only else self.trunk.embed_dim, self.vtp_config.vtp_model.text_encoder.embed_dim, bias=False) |
| 218 | else: |
| 219 | self.proj = None |
| 220 | |
| 221 | if self.vtp_config.training.train_dinov2: |
| 222 | self.dino_head = DINOHead( |
| 223 | in_dim=effective_embed_dim if not self.vtp_config.vtp_model.vision_encoder.bottleneck_ae_only else self.trunk.embed_dim, |
| 224 | out_dim=self.vtp_config.vtp_model.dino_head.out_dim, |
| 225 | nlayers=self.vtp_config.vtp_model.dino_head.nlayers, |
| 226 | hidden_dim=self.vtp_config.vtp_model.dino_head.hidden_dim, |
| 227 | bottleneck_dim=self.vtp_config.vtp_model.dino_head.bottleneck_dim, |
| 228 | ) |
| 229 | else: |
| 230 | self.dino_head = None |
| 231 | |
| 232 | if self.vtp_config.training.train_reconstruction: |
| 233 | decoder_kwargs = self.vtp_config.vtp_model.pixel_decoder |
| 234 | if self.vtp_config.vtp_model.pixel_decoder.model_type == 'dinov3': |
| 235 | decoder_kwargs["in_chans"] = effective_embed_dim |
| 236 | self.pixel_decoder = DinoV3PixelDecoder( |
| 237 | **decoder_kwargs |
| 238 | ) |
| 239 | else: |
| 240 | raise ValueError(f"Unsupported pixel decoder: {self.vtp_config.vtp_model.pixel_decoder.model_type}") |
| 241 | else: |
| 242 | self.pixel_decoder = None |
| 243 | |
| 244 | if self.vtp_config.training.train_dinov2: |
| 245 | self.enable_teacher = True |
| 246 | with torch.no_grad(): |
| 247 | self.teacher_trunk = copy.deepcopy(self.trunk) |
| 248 | if self.vtp_config.training.train_clip: |
| 249 | self.teacher_proj = nn.Linear( |
| 250 | effective_embed_dim if not self.vtp_config.vtp_model.vision_encoder.bottleneck_ae_only else self.trunk.embed_dim, |
| 251 | self.vtp_config.vtp_model.text_encoder.embed_dim, bias=False |
no test coverage detected