MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / ViTModel

Class ViTModel

SwissArmyTransformer/sat/model/official/vit_model.py:105–126  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

103
104
105class 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

Callers 2

transform_param.pyFile · 0.90
transform_param.pyFile · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected