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