(
self,
pos_type='rope100',
decoder_size='large',
)
| 15 | |
| 16 | class Pi3(nn.Module, PyTorchModelHubMixin): |
| 17 | def __init__( |
| 18 | self, |
| 19 | pos_type='rope100', |
| 20 | decoder_size='large', |
| 21 | ): |
| 22 | super().__init__() |
| 23 | |
| 24 | # ---------------------- |
| 25 | # Encoder |
| 26 | # ---------------------- |
| 27 | self.encoder = dinov2_vitl14_reg(pretrained=False) |
| 28 | self.patch_size = 14 |
| 29 | del self.encoder.mask_token |
| 30 | |
| 31 | # ---------------------- |
| 32 | # Positonal Encoding |
| 33 | # ---------------------- |
| 34 | self.pos_type = pos_type if pos_type is not None else 'none' |
| 35 | self.rope=None |
| 36 | if self.pos_type.startswith('rope'): # eg rope100 |
| 37 | if RoPE2D is None: raise ImportError("Cannot find cuRoPE2D, please install it following the README instructions") |
| 38 | freq = float(self.pos_type[len('rope'):]) |
| 39 | self.rope = RoPE2D(freq=freq) |
| 40 | self.position_getter = PositionGetter() |
| 41 | else: |
| 42 | raise NotImplementedError |
| 43 | |
| 44 | |
| 45 | # ---------------------- |
| 46 | # Decoder |
| 47 | # ---------------------- |
| 48 | enc_embed_dim = self.encoder.blocks[0].attn.qkv.in_features # 1024 |
| 49 | if decoder_size == 'small': |
| 50 | dec_embed_dim = 384 |
| 51 | dec_num_heads = 6 |
| 52 | mlp_ratio = 4 |
| 53 | dec_depth = 24 |
| 54 | elif decoder_size == 'base': |
| 55 | dec_embed_dim = 768 |
| 56 | dec_num_heads = 12 |
| 57 | mlp_ratio = 4 |
| 58 | dec_depth = 24 |
| 59 | elif decoder_size == 'large': |
| 60 | dec_embed_dim = 1024 |
| 61 | dec_num_heads = 16 |
| 62 | mlp_ratio = 4 |
| 63 | dec_depth = 36 |
| 64 | else: |
| 65 | raise NotImplementedError |
| 66 | self.decoder = nn.ModuleList([ |
| 67 | BlockRope( |
| 68 | dim=dec_embed_dim, |
| 69 | num_heads=dec_num_heads, |
| 70 | mlp_ratio=mlp_ratio, |
| 71 | qkv_bias=True, |
| 72 | proj_bias=True, |
| 73 | ffn_bias=True, |
| 74 | drop_path=0.0, |
nothing calls this directly
no test coverage detected