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