"""DETR-style transformer encoder/decoder. Positional encoding is injected into queries/keys at every attention call (not just added once at the input), following DETR -- the paper's own ablation found 4 encoder layers / 1 decoder layer to be the sweet spot before overfitting, which we reuse here. """ from __future__ import annotations import torch import torch.nn as nn def _with_pos(x: torch.Tensor, pos: torch.Tensor | None) -> torch.Tensor: return x if pos is None else x + pos class EncoderLayer(nn.Module): def __init__(self, d_model: int, nhead: int, dim_feedforward: int, dropout: float): super().__init__() self.self_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout, batch_first=True) self.ffn = nn.Sequential( nn.Linear(d_model, dim_feedforward), nn.ReLU(inplace=True), nn.Dropout(dropout), nn.Linear(dim_feedforward, d_model), ) self.norm1 = nn.LayerNorm(d_model) self.norm2 = nn.LayerNorm(d_model) self.drop1 = nn.Dropout(dropout) self.drop2 = nn.Dropout(dropout) def forward(self, src: torch.Tensor, pos: torch.Tensor) -> torch.Tensor: q = k = _with_pos(src, pos) attn_out, _ = self.self_attn(q, k, src) src = self.norm1(src + self.drop1(attn_out)) ffn_out = self.ffn(src) src = self.norm2(src + self.drop2(ffn_out)) return src class Encoder(nn.Module): def __init__(self, num_layers: int, d_model: int, nhead: int, dim_feedforward: int, dropout: float): super().__init__() self.layers = nn.ModuleList([ EncoderLayer(d_model, nhead, dim_feedforward, dropout) for _ in range(num_layers) ]) def forward(self, src: torch.Tensor, pos: torch.Tensor) -> torch.Tensor: for layer in self.layers: src = layer(src, pos) return src class DecoderLayer(nn.Module): def __init__(self, d_model: int, nhead: int, dim_feedforward: int, dropout: float): super().__init__() self.self_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout, batch_first=True) self.cross_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout, batch_first=True) self.ffn = nn.Sequential( nn.Linear(d_model, dim_feedforward), nn.ReLU(inplace=True), nn.Dropout(dropout), nn.Linear(dim_feedforward, d_model), ) self.norm1 = nn.LayerNorm(d_model) self.norm2 = nn.LayerNorm(d_model) self.norm3 = nn.LayerNorm(d_model) self.drop1 = nn.Dropout(dropout) self.drop2 = nn.Dropout(dropout) self.drop3 = nn.Dropout(dropout) def forward(self, tgt: torch.Tensor, memory: torch.Tensor, query_pos: torch.Tensor, memory_pos: torch.Tensor) -> torch.Tensor: q = k = _with_pos(tgt, query_pos) attn_out, _ = self.self_attn(q, k, tgt) tgt = self.norm1(tgt + self.drop1(attn_out)) attn_out, _ = self.cross_attn( _with_pos(tgt, query_pos), _with_pos(memory, memory_pos), memory, ) tgt = self.norm2(tgt + self.drop2(attn_out)) ffn_out = self.ffn(tgt) tgt = self.norm3(tgt + self.drop3(ffn_out)) return tgt class Decoder(nn.Module): def __init__(self, num_layers: int, d_model: int, nhead: int, dim_feedforward: int, dropout: float): super().__init__() self.layers = nn.ModuleList([ DecoderLayer(d_model, nhead, dim_feedforward, dropout) for _ in range(num_layers) ]) def forward(self, tgt: torch.Tensor, memory: torch.Tensor, query_pos: torch.Tensor, memory_pos: torch.Tensor) -> torch.Tensor: for layer in self.layers: tgt = layer(tgt, memory, query_pos, memory_pos) return tgt