microvit-s2 / microvit_hf.py
henriquequeirozcunha's picture
Upload MicroViT-S2 as proper HF model (trust_remote_code)
b8dc1cc verified
Raw
History Blame Contribute Delete
3.27 kB
"""MicroViT HuggingFace model definition (trust_remote_code=True).
MicroViTConfig + MicroViTForImageClassification — registered with AutoClasses.
Architecture: SHViT (Single-Head Vision Transformer).
Source: https://github.com/novendrastywn/MicroViT
"""
from __future__ import annotations
import torch
import torch.nn as nn
from transformers import PretrainedConfig, PreTrainedModel
from transformers.modeling_outputs import ImageClassifierOutput
from shvit import SHViT # absolute import — both files live in HF repo root
_VARIANT_CONFIGS: dict[str, dict] = {
"s1": {
"embed_dim": [128, 224, 320],
"partial_dim": [32, 48, 68],
"qk_dim": [16, 16, 16],
"depth": [2, 4, 5],
"types": ["i", "s", "s"],
"down_ops": [["subsample", 2], ["subsample", 2], [""]],
},
"s2": {
"embed_dim": [128, 308, 448],
"partial_dim": [32, 66, 96],
"qk_dim": [16, 16, 16],
"depth": [2, 4, 5],
"types": ["i", "s", "s"],
"down_ops": [["subsample", 2], ["subsample", 2], [""]],
},
"s3": {
"embed_dim": [192, 352, 448],
"partial_dim": [48, 75, 96],
"qk_dim": [16, 16, 16],
"depth": [3, 5, 5],
"types": ["i", "s", "s"],
"down_ops": [["subsample", 2], ["subsample", 2], [""]],
},
}
class MicroViTConfig(PretrainedConfig):
model_type = "microvit"
def __init__(
self,
variant: str = "s1",
embed_dim: list[int] | None = None,
partial_dim: list[int] | None = None,
qk_dim: list[int] | None = None,
depth: list[int] | None = None,
types: list[str] | None = None,
down_ops: list | None = None,
num_labels: int = 1000,
**kwargs,
):
super().__init__(num_labels=num_labels, **kwargs)
defaults = _VARIANT_CONFIGS.get(variant, _VARIANT_CONFIGS["s1"])
self.variant = variant
self.embed_dim = embed_dim if embed_dim is not None else defaults["embed_dim"]
self.partial_dim = partial_dim if partial_dim is not None else defaults["partial_dim"]
self.qk_dim = qk_dim if qk_dim is not None else defaults["qk_dim"]
self.depth = depth if depth is not None else defaults["depth"]
self.types = types if types is not None else defaults["types"]
self.down_ops = down_ops if down_ops is not None else defaults["down_ops"]
class MicroViTForImageClassification(PreTrainedModel):
config_class = MicroViTConfig
def __init__(self, config: MicroViTConfig):
super().__init__(config)
self.backbone = SHViT(
num_classes=config.num_labels,
embed_dim=config.embed_dim,
partial_dim=config.partial_dim,
qk_dim=config.qk_dim,
depth=config.depth,
types=config.types,
down_ops=config.down_ops,
)
self.post_init()
def forward(
self,
pixel_values: torch.Tensor,
labels: torch.Tensor | None = None,
) -> ImageClassifierOutput:
logits = self.backbone(pixel_values)
loss = None
if labels is not None:
loss = nn.CrossEntropyLoss()(logits, labels)
return ImageClassifierOutput(loss=loss, logits=logits)