from __future__ import annotations from copy import deepcopy from dataclasses import dataclass from typing import Any import torch from torch import Tensor, nn import torch.nn.functional as F from .boxes import inverse_sigmoid class ConvNormAct(nn.Sequential): def __init__( self, in_channels: int, out_channels: int, kernel_size: int = 1, stride: int = 1, groups: int = 1, activation: bool = True, ) -> None: padding = kernel_size // 2 layers: list[nn.Module] = [ nn.Conv2d( in_channels, out_channels, kernel_size, stride, padding, groups=groups, bias=False, ), nn.BatchNorm2d(out_channels), ] if activation: layers.append(nn.SiLU(inplace=True)) super().__init__(*layers) class GatedConvBlock(nn.Module): """Inverted residual block with a cheap learned residual gate.""" def __init__(self, channels: int, expansion: float = 2.0) -> None: super().__init__() hidden = int(channels * expansion) self.expand = ConvNormAct(channels, hidden) self.depthwise = ConvNormAct(hidden, hidden, 3, groups=hidden) self.project = ConvNormAct(hidden, channels, activation=False) self.gate = nn.Parameter(torch.zeros(1)) def forward(self, inputs: Tensor) -> Tensor: return inputs + torch.tanh(self.gate) * self.project(self.depthwise(self.expand(inputs))) class BackboneStage(nn.Sequential): def __init__(self, in_channels: int, out_channels: int, depth: int, stride: int) -> None: super().__init__( ConvNormAct(in_channels, out_channels, 3, stride=stride), *(GatedConvBlock(out_channels) for _ in range(depth)), ) class CompactBackbone(nn.Module): def __init__( self, stem_channels: int, channels: list[int], depths: list[int] ) -> None: super().__init__() if len(channels) != 4 or len(depths) != 4: raise ValueError("Backbone requires four channel and depth values") self.stem = nn.Sequential( ConvNormAct(3, stem_channels, 3, stride=2), ConvNormAct(stem_channels, stem_channels, 3, stride=2), ) stages: list[nn.Module] = [] in_channels = stem_channels for index, (out_channels, depth) in enumerate(zip(channels, depths, strict=True)): stages.append( BackboneStage(in_channels, out_channels, depth, stride=1 if index == 0 else 2) ) in_channels = out_channels self.stages = nn.ModuleList(stages) self.out_channels = channels[1:] def forward(self, images: Tensor) -> list[Tensor]: features = self.stem(images) outputs = [] for index, stage in enumerate(self.stages): features = stage(features) if index > 0: outputs.append(features) return outputs class PyramidFusion(nn.Module): def __init__(self, in_channels: list[int], hidden_dim: int, depth: int) -> None: super().__init__() self.lateral = nn.ModuleList(ConvNormAct(c, hidden_dim) for c in in_channels) self.refine = nn.ModuleList( nn.Sequential(*(GatedConvBlock(hidden_dim, expansion=1.5) for _ in range(depth))) for _ in in_channels ) def forward(self, inputs: list[Tensor]) -> list[Tensor]: projected = [layer(x) for layer, x in zip(self.lateral, inputs, strict=True)] outputs = list(projected) for index in range(len(outputs) - 2, -1, -1): outputs[index] = outputs[index] + F.interpolate( outputs[index + 1], size=outputs[index].shape[-2:], mode="nearest" ) return [block(x) for block, x in zip(self.refine, outputs, strict=True)] def sine_position_encoding( height: int, width: int, dim: int, device: torch.device, dtype: torch.dtype ) -> Tensor: if dim % 4 != 0: raise ValueError("Position encoding dimension must be divisible by four") y, x = torch.meshgrid( torch.linspace(0, 1, height, device=device, dtype=dtype), torch.linspace(0, 1, width, device=device, dtype=dtype), indexing="ij", ) frequencies = torch.arange(dim // 4, device=device, dtype=dtype) frequencies = 2.0 * torch.pi * (10000.0 ** (-frequencies / max(dim // 4, 1))) x = x.flatten()[:, None] * frequencies[None] y = y.flatten()[:, None] * frequencies[None] return torch.cat((x.sin(), x.cos(), y.sin(), y.cos()), dim=-1) class FeedForward(nn.Sequential): def __init__(self, dim: int, expansion: int = 4, dropout: float = 0.0) -> None: super().__init__( nn.Linear(dim, dim * expansion), nn.GELU(), nn.Dropout(dropout), nn.Linear(dim * expansion, dim), nn.Dropout(dropout), ) class LatentLayer(nn.Module): def __init__(self, dim: int, num_heads: int, dropout: float) -> None: super().__init__() self.norm1 = nn.LayerNorm(dim) self.attention = nn.MultiheadAttention(dim, num_heads, dropout, batch_first=True) self.norm2 = nn.LayerNorm(dim) self.ffn = FeedForward(dim, dropout=dropout) def forward(self, inputs: Tensor) -> Tensor: normalized = self.norm1(inputs) inputs = inputs + self.attention(normalized, normalized, normalized, need_weights=False)[0] return inputs + self.ffn(self.norm2(inputs)) class LatentMemory(nn.Module): """Compresses multi-scale maps into a fixed-size global reasoning memory.""" def __init__( self, dim: int, latent_count: int, pool_sizes: list[int], layers: int, num_heads: int, dropout: float, ) -> None: super().__init__() if len(pool_sizes) != 3: raise ValueError("One latent pool size is required for each pyramid level") self.pool_sizes = pool_sizes self.latents = nn.Parameter(torch.empty(latent_count, dim)) self.level_embedding = nn.Parameter(torch.empty(len(pool_sizes), dim)) self.query_norm = nn.LayerNorm(dim) self.token_norm = nn.LayerNorm(dim) self.compress = nn.MultiheadAttention(dim, num_heads, dropout, batch_first=True) self.layers = nn.ModuleList(LatentLayer(dim, num_heads, dropout) for _ in range(layers)) nn.init.normal_(self.latents, std=0.02) nn.init.normal_(self.level_embedding, std=0.02) def forward(self, features: list[Tensor]) -> Tensor: tokens = [] for level, (feature, size) in enumerate(zip(features, self.pool_sizes, strict=True)): pooled = F.adaptive_avg_pool2d(feature, (size, size)).flatten(2).transpose(1, 2) position = sine_position_encoding( size, size, feature.shape[1], feature.device, feature.dtype ) tokens.append(pooled + position[None] + self.level_embedding[level][None, None]) token_memory = self.token_norm(torch.cat(tokens, dim=1)) latents = self.latents[None].expand(features[0].shape[0], -1, -1) latents = latents + self.compress( self.query_norm(latents), token_memory, token_memory, need_weights=False )[0] for layer in self.layers: latents = layer(latents) return latents class QueryLocalSampler(nn.Module): """Samples high-resolution pyramid evidence around each evolving query box.""" def __init__(self, dim: int, num_levels: int, points: int) -> None: super().__init__() self.num_levels = num_levels self.points = points self.offsets = nn.Linear(dim, num_levels * points * 2) self.weights = nn.Linear(dim, num_levels * points) self.output = nn.Linear(dim, dim) nn.init.zeros_(self.offsets.weight) nn.init.zeros_(self.offsets.bias) nn.init.zeros_(self.weights.weight) nn.init.zeros_(self.weights.bias) def forward(self, queries: Tensor, boxes: Tensor, features: list[Tensor]) -> Tensor: batch, query_count, _ = queries.shape offsets = self.offsets(queries).view( batch, query_count, self.num_levels, self.points, 2 ) offsets = offsets.tanh() * boxes[..., None, None, 2:] * 0.5 centers = boxes[..., None, None, :2] sample_points = (centers + offsets).clamp(0.0, 1.0) weights = self.weights(queries).view( batch, query_count, self.num_levels * self.points ) weights = weights.softmax(dim=-1).view( batch, query_count, self.num_levels, self.points ) sampled_levels = [] for level, feature in enumerate(features): grid = sample_points[:, :, level] * 2.0 - 1.0 sampled = F.grid_sample( feature, grid, mode="bilinear", padding_mode="zeros", align_corners=False, ) sampled = sampled.permute(0, 2, 3, 1) sampled_levels.append(sampled) sampled_features = torch.stack(sampled_levels, dim=2) fused = (sampled_features * weights[..., None]).sum(dim=(2, 3)) return self.output(fused) class DecoderLayer(nn.Module): def __init__( self, dim: int, num_heads: int, num_levels: int, local_points: int, dropout: float ) -> None: super().__init__() self.norm1 = nn.LayerNorm(dim) self.self_attention = nn.MultiheadAttention(dim, num_heads, dropout, batch_first=True) self.norm2 = nn.LayerNorm(dim) self.global_attention = nn.MultiheadAttention(dim, num_heads, dropout, batch_first=True) self.norm3 = nn.LayerNorm(dim) self.local_sampler = QueryLocalSampler(dim, num_levels, local_points) self.norm4 = nn.LayerNorm(dim) self.ffn = FeedForward(dim, dropout=dropout) def forward( self, queries: Tensor, memory: Tensor, boxes: Tensor, features: list[Tensor] ) -> Tensor: normalized = self.norm1(queries) queries = queries + self.self_attention( normalized, normalized, normalized, need_weights=False )[0] queries = queries + self.global_attention( self.norm2(queries), memory, memory, need_weights=False )[0] queries = queries + self.local_sampler(self.norm3(queries), boxes, features) return queries + self.ffn(self.norm4(queries)) class MLP(nn.Sequential): def __init__(self, input_dim: int, hidden_dim: int, output_dim: int, layers: int) -> None: modules: list[nn.Module] = [] for index in range(layers): in_dim = input_dim if index == 0 else hidden_dim out_dim = output_dim if index == layers - 1 else hidden_dim modules.append(nn.Linear(in_dim, out_dim)) if index < layers - 1: modules.append(nn.ReLU(inplace=True)) super().__init__(*modules) class DenseAuxiliaryHead(nn.Module): def __init__(self, dim: int, num_classes: int) -> None: super().__init__() self.shared = nn.ModuleList( nn.Sequential(ConvNormAct(dim, dim, 3, groups=dim), ConvNormAct(dim, dim)) for _ in range(3) ) self.classification = nn.Conv2d(dim, num_classes, 1) self.regression = nn.Conv2d(dim, 4, 1) def forward(self, features: list[Tensor]) -> list[dict[str, Tensor]]: outputs = [] for feature, tower in zip(features, self.shared, strict=True): hidden = tower(feature) outputs.append( { "logits": self.classification(hidden), "distances": F.softplus(self.regression(hidden)), } ) return outputs @dataclass(frozen=True) class ObjectModelV1Spec: num_classes: int = 80 input_size: int = 640 stem_channels: int = 48 backbone_channels: tuple[int, int, int, int] = (64, 128, 256, 384) backbone_depths: tuple[int, int, int, int] = (2, 3, 6, 3) hidden_dim: int = 256 fpn_depth: int = 2 latent_count: int = 64 latent_pool_sizes: tuple[int, int, int] = (12, 6, 3) latent_layers: int = 2 decoder_layers: int = 6 num_queries: int = 300 num_heads: int = 8 local_points: int = 4 dropout: float = 0.0 dense_aux: bool = True class ObjectModelV1(nn.Module): """NMS-free detector with compressed global memory and local geometric sampling.""" def __init__(self, spec: ObjectModelV1Spec) -> None: super().__init__() self.spec = spec self.backbone = CompactBackbone( spec.stem_channels, list(spec.backbone_channels), list(spec.backbone_depths) ) self.neck = PyramidFusion(self.backbone.out_channels, spec.hidden_dim, spec.fpn_depth) self.memory = LatentMemory( spec.hidden_dim, spec.latent_count, list(spec.latent_pool_sizes), spec.latent_layers, spec.num_heads, spec.dropout, ) decoder_template = DecoderLayer( spec.hidden_dim, spec.num_heads, 3, spec.local_points, spec.dropout ) self.decoder = nn.ModuleList(deepcopy(decoder_template) for _ in range(spec.decoder_layers)) self.query_embedding = nn.Embedding(spec.num_queries, spec.hidden_dim) self.reference_points = nn.Embedding(spec.num_queries, 4) self.class_heads = nn.ModuleList( nn.Linear(spec.hidden_dim, spec.num_classes) for _ in range(spec.decoder_layers) ) self.box_heads = nn.ModuleList( MLP(spec.hidden_dim, spec.hidden_dim, 4, 3) for _ in range(spec.decoder_layers) ) self.dense_head = ( DenseAuxiliaryHead(spec.hidden_dim, spec.num_classes) if spec.dense_aux else None ) self._reset_parameters() def _reset_parameters(self) -> None: prior_probability = 0.01 class_bias = -torch.log(torch.tensor((1.0 - prior_probability) / prior_probability)) for head in self.class_heads: nn.init.constant_(head.bias, class_bias) nn.init.zeros_(self.reference_points.weight) with torch.no_grad(): self.reference_points.weight[:, 2:] = -2.0 for head in self.box_heads: nn.init.zeros_(head[-1].weight) nn.init.zeros_(head[-1].bias) if self.dense_head is not None: nn.init.constant_(self.dense_head.classification.bias, class_bias) nn.init.zeros_(self.dense_head.regression.weight) nn.init.constant_(self.dense_head.regression.bias, 1.0) def forward(self, images: Tensor) -> dict[str, Any]: features = self.neck(self.backbone(images)) memory = self.memory(features) batch = images.shape[0] queries = self.query_embedding.weight[None].expand(batch, -1, -1) boxes = self.reference_points.weight.sigmoid()[None].expand(batch, -1, -1) layer_outputs: list[dict[str, Tensor]] = [] for layer, class_head, box_head in zip( self.decoder, self.class_heads, self.box_heads, strict=True ): queries = layer(queries, memory, boxes, features) boxes = (inverse_sigmoid(boxes) + box_head(queries)).sigmoid() layer_outputs.append({"pred_logits": class_head(queries), "pred_boxes": boxes}) boxes = boxes.detach() if self.training else boxes output: dict[str, Any] = dict(layer_outputs[-1]) output["aux_outputs"] = layer_outputs[:-1] if self.training and self.dense_head is not None: output["dense_outputs"] = self.dense_head(features) return output def build_model(config: dict[str, Any]) -> ObjectModelV1: model_config = config.get("model", config) fields = ObjectModelV1Spec.__dataclass_fields__ unknown = set(model_config) - set(fields) if unknown: raise ValueError(f"Unknown model configuration keys: {sorted(unknown)}") values = dict(model_config) for key in ("backbone_channels", "backbone_depths", "latent_pool_sizes"): if key in values: values[key] = tuple(values[key]) return ObjectModelV1(ObjectModelV1Spec(**values))