Download src/decoder.py from shubhexists/asr: direct link, hf CLI and curl.
- Browser
- Download file 4.05 kB
-
https://huggingface.co/shubhexists/asr/resolve/main/src/decoder.py
- Command line
-
hf download hf://shubhexists/asr/src/decoder.py
-
curl -L -o decoder.py https://huggingface.co/shubhexists/asr/resolve/main/src/decoder.py
4.05 kB
| """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) | |
| 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 | |