""" Hugging Face compatible model for Trendyol DinoV2.1 (GeM pooling). Checkpoint source: ray-dinov2-full_catalog_1000_20-pfc-gem-mlp-run_12 / epoch=09 """ from __future__ import annotations from typing import Optional import torch import torch.nn as nn import torch.nn.functional as F from transformers import PretrainedConfig, PreTrainedModel from transformers.modeling_outputs import BaseModelOutput class TrendyolDinoV21Config(PretrainedConfig): model_type = "trendyol_dinov2_v21" def __init__( self, embedding_dim: int = 256, input_size: int = 224, backbone_name: str = "dinov2_vitb14", in_features: int = 768, gem_p: float = 2.9, dropout: float = 0.3, pad_color: int = 255, **kwargs, ): super().__init__(**kwargs) self.embedding_dim = embedding_dim self.input_size = input_size self.backbone_name = backbone_name self.in_features = in_features self.gem_p = gem_p self.dropout = dropout self.pad_color = pad_color self.hidden_size = embedding_dim class GeM(nn.Module): def __init__(self, p: float = 2.9, eps: float = 1e-6): super().__init__() self.register_buffer("p", torch.tensor(float(p))) self.eps = eps def forward(self, x: torch.Tensor) -> torch.Tensor: x = x.clamp(min=self.eps).pow(self.p) x = F.adaptive_avg_pool2d(x, (1, 1)) x = x.pow(1.0 / self.p) return x.squeeze(-1).squeeze(-1) class DinoV2Backbone(nn.Module): """Matches training ``Backbone`` wrapper: ``self.model = hub dinov2``.""" def __init__(self, backbone_name: str = "dinov2_vitb14"): super().__init__() try: self.model = torch.hub.load("facebookresearch/dinov2", backbone_name) except Exception as exc: # noqa: BLE001 raise RuntimeError(f"Failed to load DinoV2 backbone: {exc}") from exc self.model.requires_grad_(False) def forward(self, x: torch.Tensor) -> torch.Tensor: features = self.model.get_intermediate_layers( x, return_class_token=True, reshape=True ) return features[0][0] class TYGeMDinoV2(nn.Module): """Core GeM retrieval trunk used in run_12 (embeddings only).""" def __init__(self, config: TrendyolDinoV21Config): super().__init__() self.config = config self.backbone = DinoV2Backbone(config.backbone_name) self.pooling = GeM(p=config.gem_p) self.feature = nn.Sequential( nn.Linear(config.in_features, config.embedding_dim, bias=False), nn.BatchNorm1d(config.embedding_dim), nn.Dropout(config.dropout), ) def forward(self, pixel_values: torch.Tensor) -> torch.Tensor: feats = self.backbone(pixel_values) feats = self.pooling(feats) feats = self.feature(feats) return F.normalize(feats, p=2, dim=1) class TrendyolDinoV21Model(PreTrainedModel): config_class = TrendyolDinoV21Config base_model_prefix = "model" def __init__(self, config: TrendyolDinoV21Config): super().__init__(config) self.model = TYGeMDinoV2(config) self.post_init() @classmethod def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs): # torch.hub DinoV2 backbone is incompatible with meta-device init. kwargs.setdefault("low_cpu_mem_usage", False) return super().from_pretrained(pretrained_model_name_or_path, *model_args, **kwargs) def forward( self, pixel_values: Optional[torch.Tensor] = None, return_dict: Optional[bool] = None, **kwargs, ): return_dict = return_dict if return_dict is not None else self.config.use_return_dict if pixel_values is None: raise ValueError("pixel_values cannot be None") embeddings = self.model(pixel_values) if not return_dict: return (embeddings,) return BaseModelOutput( last_hidden_state=embeddings, hidden_states=None, attentions=None, ) def get_embeddings(self, pixel_values: torch.Tensor) -> torch.Tensor: with torch.no_grad(): return self.forward(pixel_values, return_dict=True).last_hidden_state TrendyolDinoV21Config.register_for_auto_class() TrendyolDinoV21Model.register_for_auto_class("AutoModel")