# Copyright 2026 SARMAE Authors and The HuggingFace Inc. team. """Self-contained SARMAE model and config for trust_remote_code loading.""" from __future__ import annotations from functools import partial from typing import Optional import numpy as np import torch from timm.models.vision_transformer import Block, PatchEmbed from torch import nn from transformers.configuration_utils import PretrainedConfig as PreTrainedConfig from transformers.modeling_outputs import BaseModelOutputWithPooling, ImageClassifierOutput from transformers.modeling_utils import PreTrainedModel from transformers.processing_utils import Unpack from transformers.utils import TransformersKwargs, logging logger = logging.get_logger(__name__) IMAGENET_MEAN = [0.485, 0.456, 0.406] IMAGENET_STD = [0.229, 0.224, 0.225] def get_2d_sincos_pos_embed(embed_dim: int, grid_size: int, cls_token: bool = False) -> np.ndarray: grid_h = np.arange(grid_size, dtype=np.float32) grid_w = np.arange(grid_size, dtype=np.float32) grid = np.meshgrid(grid_w, grid_h) grid = np.stack(grid, axis=0).reshape([2, 1, grid_size, grid_size]) pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid) if cls_token: pos_embed = np.concatenate([np.zeros([1, embed_dim]), pos_embed], axis=0) return pos_embed def get_2d_sincos_pos_embed_from_grid(embed_dim: int, grid: np.ndarray) -> np.ndarray: emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0]) emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1]) return np.concatenate([emb_h, emb_w], axis=1) def get_1d_sincos_pos_embed_from_grid(embed_dim: int, pos: np.ndarray) -> np.ndarray: omega = np.arange(embed_dim // 2, dtype=np.float32) omega /= embed_dim / 2.0 omega = 1.0 / 10000**omega pos = pos.reshape(-1) out = np.einsum("m,d->md", pos, omega) return np.concatenate([np.sin(out), np.cos(out)], axis=1) class SarmaeConfig(PreTrainedConfig): model_type = "sarmae" def __init__( self, hidden_size: int = 768, num_hidden_layers: int = 12, num_attention_heads: int = 12, intermediate_size: int | None = None, hidden_act: str = "gelu", hidden_dropout_prob: float = 0.0, attention_probs_dropout_prob: float = 0.0, initializer_range: float = 0.02, layer_norm_eps: float = 1e-6, image_size: int = 224, patch_size: int = 16, num_channels: int = 3, qkv_bias: bool = True, mlp_ratio: float = 4.0, global_pool: bool = True, repeat_grayscale_channels: bool = True, checkpoint_stage: str = "pretrain", image_mean: list[float] | None = None, image_std: list[float] | None = None, num_labels: int = 0, **kwargs, ): super().__init__(**kwargs) self.hidden_size = hidden_size self.num_hidden_layers = num_hidden_layers self.num_attention_heads = num_attention_heads self.hidden_act = hidden_act self.hidden_dropout_prob = hidden_dropout_prob self.attention_probs_dropout_prob = attention_probs_dropout_prob self.initializer_range = initializer_range self.layer_norm_eps = layer_norm_eps self.image_size = image_size self.patch_size = patch_size self.num_channels = num_channels self.qkv_bias = qkv_bias self.mlp_ratio = mlp_ratio self.global_pool = global_pool self.repeat_grayscale_channels = repeat_grayscale_channels self.checkpoint_stage = checkpoint_stage self.num_labels = num_labels self.intermediate_size = int(hidden_size * mlp_ratio) if intermediate_size is None else intermediate_size self.image_mean = image_mean if image_mean is not None else IMAGENET_MEAN self.image_std = image_std if image_std is not None else IMAGENET_STD class SarmaePreTrainedModel(PreTrainedModel): config_class = SarmaeConfig config: SarmaeConfig base_model_prefix = "sarmae" main_input_name = "pixel_values" input_modalities = ("image",) supports_gradient_checkpointing = True _no_split_modules = ["Block"] class SarmaeModel(SarmaePreTrainedModel): def __init__(self, config: SarmaeConfig, add_pooling_layer: bool = True): super().__init__(config) self.config = config self.add_pooling_layer = add_pooling_layer image_size = config.image_size if isinstance(config.image_size, int) else config.image_size[0] norm_layer = partial(nn.LayerNorm, eps=config.layer_norm_eps) self.patch_embed = PatchEmbed(image_size, config.patch_size, config.num_channels, config.hidden_size) self.num_patches = self.patch_embed.num_patches self.cls_token = nn.Parameter(torch.zeros(1, 1, config.hidden_size)) self.pos_embed = nn.Parameter(torch.zeros(1, self.num_patches + 1, config.hidden_size)) pos_embed = get_2d_sincos_pos_embed(self.pos_embed.shape[-1], int(self.num_patches**0.5), cls_token=True) self.pos_embed.data.copy_(torch.from_numpy(pos_embed).float().unsqueeze(0)) self.blocks = nn.ModuleList([ Block(config.hidden_size, config.num_attention_heads, config.mlp_ratio, qkv_bias=config.qkv_bias, norm_layer=norm_layer) for _ in range(config.num_hidden_layers) ]) self.global_pool = config.global_pool if self.global_pool: self.fc_norm = norm_layer(config.hidden_size) self.norm = None else: self.fc_norm = None self.norm = norm_layer(config.hidden_size) self.post_init() def forward( self, pixel_values: Optional[torch.Tensor] = None, return_dict: Optional[bool] = None, **kwargs: Unpack[TransformersKwargs], ) -> BaseModelOutputWithPooling: if pixel_values is None: raise ValueError("You must specify `pixel_values`") pixel_values = pixel_values.to(dtype=self.dtype) if return_dict is None: return_dict = self.config.use_return_dict batch_size = pixel_values.shape[0] patch_tokens = self.patch_embed(pixel_values) cls_tokens = self.cls_token.expand(batch_size, -1, -1) hidden_states = torch.cat((cls_tokens, patch_tokens), dim=1) + self.pos_embed for block in self.blocks: hidden_states = block(hidden_states) if self.global_pool: pooled_output = self.fc_norm(hidden_states[:, 1:, :].mean(dim=1)) else: hidden_states = self.norm(hidden_states) pooled_output = hidden_states[:, 0] if not self.add_pooling_layer: pooled_output = None if not return_dict: return (hidden_states, pooled_output) return BaseModelOutputWithPooling(last_hidden_state=hidden_states, pooler_output=pooled_output) class SarmaeForImageClassification(SarmaePreTrainedModel): def __init__(self, config: SarmaeConfig): super().__init__(config) self.sarmae = SarmaeModel(config, add_pooling_layer=True) self.classifier = nn.Linear(config.hidden_size, config.num_labels) if config.num_labels > 0 else nn.Identity() self.post_init() def forward( self, pixel_values: Optional[torch.Tensor] = None, labels: Optional[torch.Tensor] = None, return_dict: Optional[bool] = None, **kwargs: Unpack[TransformersKwargs], ) -> ImageClassifierOutput: outputs = self.sarmae(pixel_values=pixel_values, return_dict=True, **kwargs) logits = self.classifier(outputs.pooler_output) loss = None if labels is not None: loss = self.loss_function(labels, logits, self.config, **kwargs) if not return_dict: output = (logits,) + outputs[1:] return ((loss,) + output) if loss is not None else output return ImageClassifierOutput(loss=loss, logits=logits, hidden_states=outputs.hidden_states, attentions=outputs.attentions) __all__ = [ "SarmaeConfig", "SarmaeForImageClassification", "SarmaeModel", "SarmaePreTrainedModel", ]