first commit
This commit is contained in:
55
models/laneformer.py
Normal file
55
models/laneformer.py
Normal file
@@ -0,0 +1,55 @@
|
||||
"""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)
|
||||
Reference in New Issue
Block a user