"""LaneFormer-CUSTOM Phase 1: backbone -> PE -> transformer encoder/decoder -> heads. Reasoning/verification module (feature-correction + confidence scoring) is Phase 2 -- this assembly is the minimal architecture needed to prove the core trains. """ from __future__ import annotations import torch import torch.nn as nn from models.backbone import Backbone from models.positional_encoding import PositionEmbedding2D from models.transformer import Encoder, Decoder from models.head import PredictionHeads class LaneFormer(nn.Module): def __init__( self, backbone_name: str = "resnet34", pretrained: bool = True, d_model: int = 128, max_lanes: int = 8, encoder_layers: int = 4, decoder_layers: int = 1, nhead: int = 8, ffn_dim: int = 512, dropout: float = 0.1, ): super().__init__() self.backbone = Backbone(backbone_name, pretrained=pretrained, out_channels=d_model) self.pos_embed = PositionEmbedding2D(d_model) self.encoder = Encoder(encoder_layers, d_model, nhead, ffn_dim, dropout) self.decoder = Decoder(decoder_layers, d_model, nhead, ffn_dim, dropout) self.query_embed = nn.Embedding(max_lanes, d_model) self.heads = PredictionHeads(d_model) self.max_lanes = max_lanes self.d_model = d_model def forward(self, images: torch.Tensor) -> dict: B = images.shape[0] feat = self.backbone(images) # (B, C, H, W) _, C, H, W = feat.shape src = feat.flatten(2).permute(0, 2, 1) # (B, H*W, C) pos = self.pos_embed(H, W, images.device) # (H*W, C) pos = pos.unsqueeze(0).expand(B, -1, -1) # (B, H*W, C) memory = self.encoder(src, pos) # (B, H*W, C) query_pos = self.query_embed.weight.unsqueeze(0).expand(B, -1, -1) # (B, max_lanes, C) tgt = torch.zeros_like(query_pos) decoded = self.decoder(tgt, memory, query_pos, pos) # (B, max_lanes, C) return self.heads(decoded)