import torch import torch.nn as nn import torch.nn.functional as F from dataclasses import dataclass from typing import Optional, Tuple, List from transformers import PretrainedConfig, PreTrainedModel from transformers.modeling_outputs import ModelOutput # ---------------------------------------------------------------------- # Custom Output Dataclass # ---------------------------------------------------------------------- @dataclass class ConvMixerOutput(ModelOutput): """ Output type for ConvMixerForImageClassification. Args: loss (`torch.FloatTensor`, *optional*): Classification loss. logits (`torch.FloatTensor`): Classification logits (before softmax). last_hidden_state (`torch.FloatTensor`): Sequence of spatial features (flattened), shape (batch_size, num_patches, dim). pooler_output (`torch.FloatTensor`): Attention‑pooled representation, shape (batch_size, dim). hidden_states (`tuple(torch.FloatTensor)`, *optional*): Hidden states from each block (stem + convmixer blocks). attentions (`tuple(torch.FloatTensor)`, *optional*): Attention weights from the pooling layer (if `output_attentions=True`). """ loss: Optional[torch.FloatTensor] = None logits: Optional[torch.FloatTensor] = None last_hidden_state: Optional[torch.FloatTensor] = None pooler_output: Optional[torch.FloatTensor] = None hidden_states: Optional[Tuple[torch.FloatTensor, ...]] = None attentions: Optional[Tuple[torch.FloatTensor, ...]] = None # ---------------------------------------------------------------------- # Custom Configuration # ---------------------------------------------------------------------- class ConvMixerConfig(PretrainedConfig): """ Configuration class for ConvMixer models. Args: dim (`int`, *optional*, defaults to 256): Embedding dimension throughout the network. depth (`int`, *optional*, defaults to 8): Number of ConvMixer blocks. kernel_size (`int`, *optional*, defaults to 5): Kernel size of depthwise convolutions. patch_size (`int`, *optional*, defaults to 2): Stem convolution stride / patch size. num_classes (`int`, *optional*, defaults to 1000): Number of classes for classification head. attn_pool_heads (`int`, *optional*, defaults to 8): Number of attention heads in the pooling layer. attn_pool_mlp_ratio (`float`, *optional*, defaults to 4.0): MLP hidden ratio in the pooling layer. dropout (`float`, *optional*, defaults to 0.0): Dropout rate applied in the pooling MLP. """ model_type = "convmixer" def __init__( self, dim: int = 768, depth: int = 12, kernel_size: int = 9, patch_size: int = 16, num_classes: int = 1000, attn_pool_heads: int = 16, attn_pool_mlp_ratio: float = 4.0, dropout: float = 0.0, **kwargs, ): super().__init__(**kwargs) self.dim = dim self.depth = depth self.kernel_size = kernel_size self.patch_size = patch_size self.num_classes = num_classes self.attn_pool_heads = attn_pool_heads self.attn_pool_mlp_ratio = attn_pool_mlp_ratio self.dropout = dropout # ---------------------------------------------------------------------- # Core ConvMixer Components (unchanged except removed classification head) # ---------------------------------------------------------------------- class Residual(nn.Module): """Residual wrapper used in ConvMixer.""" def __init__(self, fn): super().__init__() self.fn = fn def forward(self, x): return self.fn(x) + x class AttentionPooling(nn.Module): """ Multi‑head attention pooling that aggregates a spatial feature map into a single vector. Optionally returns attention weights. """ def __init__( self, dim: int, num_heads: int = 16, mlp_ratio: float = 4.0, dropout: float = 0.0, ): super().__init__() self.dim = dim self.num_heads = num_heads self.probe = nn.Parameter(torch.randn(1, 1, dim)) self.attention = nn.MultiheadAttention( embed_dim=dim, num_heads=num_heads, batch_first=True, dropout=dropout, ) self.layernorm = nn.LayerNorm(dim) mlp_hidden = int(dim * mlp_ratio) self.mlp = nn.Sequential( nn.Linear(dim, mlp_hidden), nn.GELU(), nn.Dropout(dropout), nn.Linear(mlp_hidden, dim), nn.Dropout(dropout), ) def forward(self, x: torch.Tensor, output_attentions: bool = False): # x shape: (B, dim, H, W) B, C, H, W = x.shape x = x.flatten(2).transpose(1, 2) # (B, L, C) probe = self.probe.expand(B, -1, -1) # (B, 1, C) attn_out, attn_weights = self.attention(probe, x, x) # (B, 1, C), (B, 1, L) residual = attn_out attn_out = self.layernorm(attn_out) attn_out = residual + self.mlp(attn_out) pooled = attn_out[:, 0] # (B, C) if output_attentions: return pooled, attn_weights return pooled, None class ConvMixerWithAttnPool(nn.Module): """ ConvMixer backbone with multi‑head attention pooling. Returns pooled representation, spatial features, and optional hidden states / attention weights. """ def __init__( self, dim: int, depth: int, kernel_size: int = 9, patch_size: int = 16, attn_pool_heads: int = 16, attn_pool_mlp_ratio: float = 4.0, dropout: float = 0.0, ): super().__init__() self.dim = dim self.depth = depth # Stem self.stem = nn.Sequential( nn.Conv2d(3, dim, kernel_size=patch_size, stride=patch_size), nn.GELU(), nn.BatchNorm2d(dim), ) # ConvMixer blocks self.blocks = nn.ModuleList([ nn.Sequential( Residual( nn.Sequential( nn.Conv2d(dim, dim, kernel_size, groups=dim, padding="same"), nn.GELU(), nn.BatchNorm2d(dim), ) ), nn.Conv2d(dim, dim, kernel_size=1), nn.GELU(), nn.BatchNorm2d(dim), ) for _ in range(depth) ]) # Attention pooling self.pool = AttentionPooling( dim=dim, num_heads=attn_pool_heads, mlp_ratio=attn_pool_mlp_ratio, dropout=dropout, ) def forward( self, x: torch.Tensor, output_hidden_states: bool = False, output_attentions: bool = False, ): """ Returns: - pooled: (B, dim) attention‑pooled vector - spatial_features: (B, L, dim) flattened spatial map before pooling - hidden_states: tuple of intermediate feature maps (if requested) - attentions: attention weights from pooling layer (if requested) """ hidden_states = () if output_hidden_states else None attentions = None # Stem x = self.stem(x) if output_hidden_states: hidden_states += (x,) # ConvMixer blocks for blk in self.blocks: x = blk(x) if output_hidden_states: hidden_states += (x,) # Store spatial features before pooling spatial_features = x.flatten(2).transpose(1, 2) # (B, L, C) # Attention pooling pooled, attn_weights = self.pool(x, output_attentions=output_attentions) if output_attentions: attentions = attn_weights return pooled, spatial_features, hidden_states, attentions # ---------------------------------------------------------------------- # Base ConvMixer Model (backbone only) # ---------------------------------------------------------------------- class ConvMixerModel(PreTrainedModel): """ Bare ConvMixer model outputting raw features and optional hidden states/attentions. This model inherits from [`PreTrainedModel`]. Check the superclass documentation for the generic methods the library implements for all its model (such as downloading or saving). It wraps the `ConvMixerWithAttnPool` backbone to be compatible with the Transformers API. """ config_class = ConvMixerConfig base_model_prefix = "convmixer" def __init__(self, config: ConvMixerConfig): super().__init__(config) self.backbone = ConvMixerWithAttnPool( dim=config.dim, depth=config.depth, kernel_size=config.kernel_size, patch_size=config.patch_size, attn_pool_heads=config.attn_pool_heads, attn_pool_mlp_ratio=config.attn_pool_mlp_ratio, dropout=config.dropout, ) # Initialize weights and apply final processing self.post_init() def forward( self, pixel_values: torch.FloatTensor, output_hidden_states: Optional[bool] = None, output_attentions: Optional[bool] = None, return_dict: Optional[bool] = None, ): output_hidden_states = output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions return_dict = return_dict if return_dict is not None else self.config.use_return_dict pooled, spatial_features, hidden_states, attentions = self.backbone( pixel_values, output_hidden_states=output_hidden_states, output_attentions=output_attentions, ) if not return_dict: return (spatial_features, pooled, hidden_states, attentions) return ConvMixerOutput( last_hidden_state=spatial_features, pooler_output=pooled, hidden_states=hidden_states, attentions=attentions, ) # ---------------------------------------------------------------------- # ConvMixer for Image Classification (Trainer‑compatible) # ---------------------------------------------------------------------- class ConvMixerForImageClassification(PreTrainedModel): """ ConvMixer model with an image classification head on top (a linear layer on top of the attention‑pooled output), e.g. for ImageNet. This model inherits from [`PreTrainedModel`]. Check the superclass documentation for the generic methods the library implements for all its model (such as downloading or saving). Args: config ([`ConvMixerConfig`]): Model configuration class with all the parameters of the model. """ config_class = ConvMixerConfig base_model_prefix = "convmixer" def __init__(self, config: ConvMixerConfig): super().__init__(config) self.backbone = ConvMixerWithAttnPool( dim=config.dim, depth=config.depth, kernel_size=config.kernel_size, patch_size=config.patch_size, attn_pool_heads=config.attn_pool_heads, attn_pool_mlp_ratio=config.attn_pool_mlp_ratio, dropout=config.dropout, ) self.classifier = nn.Linear(config.dim, config.num_classes) if config.num_classes > 0 else nn.Identity() self.loss_fn = nn.CrossEntropyLoss() if config.num_classes > 0 else None # Initialize weights and apply final processing self.post_init() def forward( self, pixel_values: torch.FloatTensor, labels: Optional[torch.LongTensor] = None, output_hidden_states: Optional[bool] = None, output_attentions: Optional[bool] = None, return_dict: Optional[bool] = None, ) -> ConvMixerOutput: r""" labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*): Labels for computing the image classification loss. Indices must be in `[0, ..., config.num_classes - 1]`. """ output_hidden_states = output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions return_dict = return_dict if return_dict is not None else self.config.use_return_dict pooled, spatial_features, hidden_states, attentions = self.backbone( pixel_values, output_hidden_states=output_hidden_states, output_attentions=output_attentions, ) logits = self.classifier(pooled) loss = None if labels is not None and self.loss_fn is not None: loss = self.loss_fn(logits, labels) if not return_dict: output = (logits, spatial_features, pooled, hidden_states, attentions) return ((loss,) + output) if loss is not None else output return ConvMixerOutput( loss=loss, logits=logits, last_hidden_state=spatial_features, pooler_output=pooled, hidden_states=hidden_states, attentions=attentions, ) # ---------------------------------------------------------------------- # Optional: Register models with auto classes for easy loading # ---------------------------------------------------------------------- ConvMixerConfig.register_for_auto_class() ConvMixerModel.register_for_auto_class("AutoModel") ConvMixerForImageClassification.register_for_auto_class("AutoModelForImageClassification")