| 103 | |
| 104 | |
| 105 | class ViTModel(BaseModel): |
| 106 | def __init__(self, args, transformer=None, **kwargs): |
| 107 | property = ViTProperty(args.image_size, args.patch_size, args.pre_len, args.post_len) |
| 108 | args.max_sequence_length = property.pre_len + property.num_patches + property.post_len |
| 109 | if 'activation_func' not in kwargs: |
| 110 | kwargs['activation_func'] = gelu |
| 111 | super().__init__(args, transformer=transformer, **kwargs) |
| 112 | self.transformer.property = property |
| 113 | self.add_mixin("patch_embedding", ImagePatchEmbeddingMixin(args.in_channels, args.hidden_size, property)) |
| 114 | self.add_mixin("pos_embedding", InterpolatedPositionEmbeddingMixin()) |
| 115 | self.add_mixin("cls", ClsMixin(args.hidden_size, args.num_classes)) |
| 116 | |
| 117 | @classmethod |
| 118 | def add_model_specific_args(cls, parser): |
| 119 | group = parser.add_argument_group('ViT', 'ViT Configurations') |
| 120 | group.add_argument('--image-size', nargs='+', type=int, default=[224, 224]) |
| 121 | group.add_argument('--pre-len', type=int, default=1) # [cls] by default |
| 122 | group.add_argument('--post-len', type=int, default=0) # empty by default, but sometimes with special tokens, such as [det] in yolos. |
| 123 | group.add_argument('--in-channels', type=int, default=3) |
| 124 | group.add_argument('--num-classes', type=int, default=21843) |
| 125 | group.add_argument('--patch-size', type=int, default=16) |
| 126 | return parser |
| 127 | |
| 128 |
no outgoing calls
no test coverage detected