"""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