94 lines
3.7 KiB
Python
94 lines
3.7 KiB
Python
"""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
|