import torch import torch.nn as nn import torch.nn.functional as F from transformers import PreTrainedModel from transformers.modeling_outputs import ImageClassifierOutput try: from .configuration_squeezenet import SqueezeNetConfig except ImportError: from configuration_squeezenet import SqueezeNetConfig class FP32LayerNorm2d(nn.GroupNorm): def __init__(self, num_channels): super().__init__(1, num_channels) def forward(self, x): input_dtype = x.dtype with torch.autocast(device_type=x.device.type, enabled=False): normalized = super().forward(x.float()) return normalized.to(dtype=input_dtype) class SwiGLU(nn.Module): def __init__(self, channels): super().__init__() self.proj = nn.Conv2d(channels, channels * 2, kernel_size=1, bias=True) self.residual_scale = nn.Parameter(torch.tensor(0.1)) def forward(self, x): gate, value = self.proj(x).chunk(2, dim=1) gated = F.silu(gate) * value return x + self.residual_scale * gated class FireModule(nn.Module): def __init__(self, in_channels, squeeze_channels, expand1x1_channels, expand3x3_channels, dropout=0.0): super().__init__() self.squeeze = nn.Conv2d(in_channels, squeeze_channels, kernel_size=1, bias=False) self.squeeze_norm = FP32LayerNorm2d(squeeze_channels) self.squeeze_act = SwiGLU(squeeze_channels) self.expand1x1 = nn.Conv2d(squeeze_channels, expand1x1_channels, kernel_size=1, bias=False) self.expand1x1_norm = FP32LayerNorm2d(expand1x1_channels) self.expand1x1_act = SwiGLU(expand1x1_channels) self.expand3x3 = nn.Conv2d(squeeze_channels, expand3x3_channels, kernel_size=3, padding=1, bias=False) self.expand3x3_norm = FP32LayerNorm2d(expand3x3_channels) self.expand3x3_act = SwiGLU(expand3x3_channels) self.dropout = nn.Dropout2d(p=dropout) if dropout > 0 else nn.Identity() def forward(self, x): x = self.squeeze_act(self.squeeze_norm(self.squeeze(x))) e1 = self.expand1x1_act(self.expand1x1_norm(self.expand1x1(x))) e3 = self.expand3x3_act(self.expand3x3_norm(self.expand3x3(x))) return self.dropout(torch.cat([e1, e3], dim=1)) class SqueezeNetForImageClassification(PreTrainedModel): config_class = SqueezeNetConfig def __init__(self, config: SqueezeNetConfig): super().__init__(config) self.num_classes = config.num_classes self.conv1 = nn.Conv2d(3, 96, kernel_size=7, stride=2, padding=3, bias=False) self.norm1 = FP32LayerNorm2d(96) self.act1 = SwiGLU(96) self.pool1 = nn.MaxPool2d(kernel_size=3, stride=2, padding=1) self.fire2 = FireModule(96, 16, 64, 64, config.fire_dropout) self.fire3 = FireModule(128, 16, 64, 64, config.fire_dropout) self.fire4 = FireModule(128, 32, 128, 128, config.fire_dropout) self.pool2 = nn.MaxPool2d(kernel_size=3, stride=2, padding=1) self.fire5 = FireModule(256, 32, 128, 128, config.fire_dropout) self.fire6 = FireModule(256, 48, 192, 192, config.fire_dropout) self.fire7 = FireModule(384, 48, 192, 192, config.fire_dropout) self.fire8 = FireModule(384, 64, 256, 256, config.fire_dropout) self.pool3 = nn.MaxPool2d(kernel_size=3, stride=2, padding=1) self.fire9 = FireModule(512, 64, 256, 256, config.fire_dropout) self.dropout = nn.Dropout2d(p=config.final_dropout) self.conv10 = nn.Conv2d(512, self.num_classes, kernel_size=1, bias=True) self.gap = nn.AdaptiveAvgPool2d((1, 1)) self.post_init() def _init_weights(self, module): if isinstance(module, nn.Conv2d): nn.init.kaiming_normal_(module.weight, mode="fan_out", nonlinearity="relu") if module.bias is not None: nn.init.zeros_(module.bias) elif isinstance(module, nn.GroupNorm): nn.init.ones_(module.weight) nn.init.zeros_(module.bias) if module == self.conv10: nn.init.normal_(module.weight, mean=0.0, std=0.001) nn.init.zeros_(module.bias) def forward(self, pixel_values: torch.Tensor, labels: torch.Tensor | None = None, return_dict: bool | None = None): return_dict = return_dict if return_dict is not None else self.config.use_return_dict x = self.pool1(self.act1(self.norm1(self.conv1(pixel_values)))) x = self.fire2(x) x = self.fire3(x) x = self.fire4(x) x = self.pool2(x) x = self.fire5(x) x = self.fire6(x) x = self.fire7(x) x = self.fire8(x) x = self.pool3(x) x = self.fire9(x) x = self.conv10(self.dropout(x)) x = self.gap(x) logits = torch.flatten(x, 1) loss = None if labels is not None: loss_fct = nn.CrossEntropyLoss() loss = loss_fct(logits.view(-1, self.num_classes), labels.view(-1)) if not return_dict: output = (logits,) return ((loss,) + output) if loss is not None else output return ImageClassifierOutput( loss=loss, logits=logits, )