56 lines
2.0 KiB
Python
56 lines
2.0 KiB
Python
"""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)
|