File size: 4,047 Bytes
ce3c8df | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 | """Standard Transformer decoder for the attention branch of the hybrid CTC/attention model."""
import math
import torch
import torch.nn as nn
from src.scaling import PositionalEncoding
from src.zipformer import make_pad_mask
class AttentionDecoder(nn.Module):
def __init__(
self,
vocab_size: int,
d_model: int = 256,
nhead: int = 4,
d_ff: int = 1024,
num_layers: int = 4,
dropout: float = 0.1,
):
super().__init__()
self.d_model = d_model
self.embed = nn.Embedding(vocab_size, d_model)
self.pos_enc = PositionalEncoding(d_model, dropout=dropout)
layer = nn.TransformerDecoderLayer(
d_model=d_model,
nhead=nhead,
dim_feedforward=d_ff,
dropout=dropout,
activation="gelu",
batch_first=True,
norm_first=True,
)
# Pre-norm layers need a final norm; the residual path is otherwise never normalized.
self.decoder = nn.TransformerDecoder(layer, num_layers=num_layers, norm=nn.LayerNorm(d_model))
self.output_proj = nn.Linear(d_model, vocab_size)
# Weight tying: share embedding and output projection weights.
self.output_proj.weight = self.embed.weight
def forward(
self,
tokens: torch.Tensor,
token_lengths: torch.Tensor,
encoder_out: torch.Tensor,
encoder_lengths: torch.Tensor,
) -> torch.Tensor:
"""
tokens: (B, U) input tokens, teacher-forced (already includes BOS, excludes final EOS target)
token_lengths: (B,) valid lengths of `tokens`
encoder_out: (B, T, d_model)
encoder_lengths: (B,)
returns: logits (B, U, vocab_size)
"""
u = tokens.size(1)
x = self.embed(tokens) * math.sqrt(self.d_model)
x = self.pos_enc(x)
causal_mask = nn.Transformer.generate_square_subsequent_mask(u, device=tokens.device)
tgt_pad_mask = make_pad_mask(token_lengths, u)
memory_pad_mask = make_pad_mask(encoder_lengths, encoder_out.size(1))
out = self.decoder(
tgt=x,
memory=encoder_out,
tgt_mask=causal_mask,
tgt_key_padding_mask=tgt_pad_mask,
memory_key_padding_mask=memory_pad_mask,
)
return self.output_proj(out)
@torch.no_grad()
def greedy_decode(
self,
encoder_out: torch.Tensor,
encoder_lengths: torch.Tensor,
bos_id: int,
eos_id: int,
max_len: int = 200,
):
"""Greedy autoregressive decode, one step at a time. Returns a list of
token-id lists (one per batch element), EOS excluded.
"""
device = encoder_out.device
b = encoder_out.size(0)
memory_pad_mask = make_pad_mask(encoder_lengths, encoder_out.size(1))
tokens = torch.full((b, 1), bos_id, dtype=torch.long, device=device)
finished = torch.zeros(b, dtype=torch.bool, device=device)
results = [[] for _ in range(b)]
for _ in range(max_len):
u = tokens.size(1)
x = self.embed(tokens) * math.sqrt(self.d_model)
x = self.pos_enc(x)
causal_mask = nn.Transformer.generate_square_subsequent_mask(u, device=device)
out = self.decoder(
tgt=x,
memory=encoder_out,
tgt_mask=causal_mask,
memory_key_padding_mask=memory_pad_mask,
)
logits = self.output_proj(out[:, -1, :]) # (B, vocab)
next_token = logits.argmax(dim=-1) # (B,)
for i in range(b):
if not finished[i]:
if next_token[i].item() == eos_id:
finished[i] = True
else:
results[i].append(next_token[i].item())
tokens = torch.cat([tokens, next_token.unsqueeze(1)], dim=1)
if finished.all():
break
return results
|