"""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)