| 76 | dim_features: int |
| 77 | |
| 78 | def __init__(self, backbone: str, intermediate_layers: Union[int, List[int]], dim_out: int, **deprecated_kwargs): |
| 79 | super(DINOv2Encoder, self).__init__() |
| 80 | |
| 81 | self.intermediate_layers = intermediate_layers |
| 82 | |
| 83 | # Load the backbone |
| 84 | self.hub_loader = getattr(importlib.import_module(".dinov2.hub.backbones", __package__), backbone) |
| 85 | self.backbone_name = backbone |
| 86 | self.backbone = self.hub_loader(pretrained=False) |
| 87 | |
| 88 | self.dim_features = self.backbone.blocks[0].attn.qkv.in_features |
| 89 | self.num_features = intermediate_layers if isinstance(intermediate_layers, int) else len(intermediate_layers) |
| 90 | |
| 91 | self.output_projections = nn.ModuleList([ |
| 92 | nn.Conv2d(in_channels=self.dim_features, out_channels=dim_out, kernel_size=1, stride=1, padding=0,) |
| 93 | for _ in range(self.num_features) |
| 94 | ]) |
| 95 | |
| 96 | self.register_buffer("image_mean", torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1)) |
| 97 | self.register_buffer("image_std", torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1)) |
| 98 | |
| 99 | def init_weights(self): |
| 100 | pretrained_backbone_state_dict = self.hub_loader(pretrained=True).state_dict() |