File size: 3,273 Bytes
b8dc1cc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
"""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)