File size: 6,214 Bytes
4ca4e4c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
import torch
import torch.nn as nn
import torch.nn.functional as F

# Causal mask: Prevent attending to future tokens
def generate_mask(sz, window=None):
    if not window:
        mask = (torch.triu(torch.ones(sz, sz)) == 1).transpose(0, 1)
    else:
        mask = (torch.triu(torch.ones(sz, sz)) - torch.triu(torch.ones(sz, sz), diagonal=window) == 1).transpose(0, 1)
    mask = mask.float().masked_fill(mask == 0, float('-inf')).masked_fill(mask == 1, float(0.0))
    return mask

####################################################################################

def scaled_dot_product_attention(q, k, v, mask=None):
    d_k = q.size(-1)
    scores = torch.matmul(q, k.transpose(-2, -1)) / torch.sqrt(torch.tensor(d_k, dtype=torch.float32))

    if mask is not None:
        scores = scores.masked_fill(mask == float('-inf'), float('-inf'))
        # scores = scores.masked_fill(mask == 0, float('-inf'))

    attn = F.softmax(scores, dim=-1)
    output = torch.matmul(attn, v)
    return output, attn


class MultiHeadAttention(nn.Module):
    def __init__(self, embed_dim, num_heads, dropout=0.1):
        super().__init__()
        assert embed_dim % num_heads == 0, "Embedding dimension must be divisible by number of heads"

        self.embed_dim = embed_dim
        self.num_heads = num_heads
        # self.head_dim = embed_dim // num_heads
        self.head_dim = embed_dim

        self.q_proj = nn.Linear(embed_dim, self.head_dim*self.num_heads)
        self.k_proj = nn.Linear(embed_dim, self.head_dim*self.num_heads)
        self.v_proj = nn.Linear(embed_dim, self.head_dim*self.num_heads)
        self.out_proj = nn.Linear(embed_dim, embed_dim)

        self.dropout = nn.Dropout(dropout)

    def forward(self, query, key, value, mask=None):
        B, T, _ = query.size()

        # Linear projections
        q = self.q_proj(query).view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
        k = self.k_proj(key).view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
        v = self.v_proj(value).view(B, T, self.num_heads, self.head_dim).transpose(1, 2)

        # Apply attention on all the projected vectors in batch
        attn_output, _ = scaled_dot_product_attention(q, k, v, mask)

        # Concatenate heads and run through final linear layer
        # attn_output = attn_output.transpose(1, 2).contiguous().view(B, T, self.embed_dim)
        attn_output = attn_output.transpose(1, 2).sum(dim=2).view(B, T, self.embed_dim)
        output = self.out_proj(attn_output)
        return output


class TransformerHead(nn.Module):
    def __init__(self, embed_dim, num_heads, ff_dim, dropout=0.1, do_norm=True):
        super().__init__()
        self.mha = MultiHeadAttention(embed_dim, num_heads, dropout)
        if do_norm:
            self.norm1 = nn.LayerNorm(embed_dim)
            self.norm2 = nn.LayerNorm(embed_dim)
        else:
            self.norm1 = nn.Identity(embed_dim)
            self.norm2 = nn.Identity(embed_dim)

        self.ffn = nn.Sequential(
            nn.Linear(embed_dim, ff_dim),
            nn.ReLU(),
            nn.Dropout(dropout),
            nn.Linear(ff_dim, embed_dim),
        )
        self.dropout = nn.Dropout(dropout)

    def forward(self, x, mask=None):
        # Multi-head attention + residual + norm
        x_norm = self.norm1(x)
        attn_out = self.mha(x_norm, x_norm, x_norm, mask)
        # x = x + self.norm1(self.dropout(attn_out))
        x = x + self.dropout(attn_out)

        # Feedforward + residual + norm
        x_norm = self.norm2(x)
        ff_out = self.ffn(x_norm)
        # x = x + self.norm2(self.dropout(ff_out))
        x = x + self.dropout(ff_out)
        return x


class SimpleHandmadeTFLayer(nn.Module):
    def __init__(self, d_model, nhead, causal=True, do_norm=True):
        super().__init__()
        self.transformer_encoder = TransformerHead(d_model, nhead, d_model, dropout=0.2, do_norm=do_norm)
        self.d_model = d_model
        self.nhead = nhead
        self.causal = causal

    def forward(self, x, mask):
        # xx = x.permute(1, 0, 2)  # (seq_len, batch_size, d_model)
        xx = x  
        if self.causal:
            xx = self.transformer_encoder(xx, mask=mask)
        else:
            xx = self.transformer_encoder(xx)
        # return xx.permute(1, 0, 2) # Already includes the skip connection
        return xx # Already includes the skip connection

    # # Finds the average required memory for select data
    # def required_activation_memory(self, x, mask, thres=0.01):
    #     B, T, _ = x.size()

    #     # Linear projections
    #     mha = self.transformer_encoder.mha
    #     q = mha.q_proj(x).view(B, T, self.nhead, self.d_model // self.nhead).transpose(1, 2)
    #     k = mha.k_proj(x).view(B, T, self.nhead, self.d_model // self.nhead).transpose(1, 2)
    #     v = mha.v_proj(x).view(B, T, self.nhead, self.d_model // self.nhead).transpose(1, 2)

    #     attn_output, attn_weights = scaled_dot_product_attention(q, k, v, mask)
    #     # This calculation is saying 'we only need the tokens where the last token attends highly with it'
    #     # It's not exactly what we are going for, so we need to think about this more
    #     return (torch.sum(attn_weights[:, 0, -1, :] > thres) / x.shape[0] * self.d_model).item()

####################################################################################

# A simple wrapper for the PyTorch transformer encoder layers
class SimplePyTorchTFLayer(nn.Module):
    def __init__(self, d_model, nhead, causal=True):
        super().__init__()
        encoder_layer = nn.TransformerEncoderLayer(d_model=d_model, nhead=nhead, dim_feedforward=d_model, dropout=0.2)
        self.transformer_encoder = nn.TransformerEncoder(encoder_layer, num_layers=1)
        self.d_model = d_model
        self.nhead = nhead
        self.causal = causal

    def forward(self, x, mask):
        xx = x.permute(1, 0, 2)  # (seq_len, batch_size, d_model)
        # xx = x
        if self.causal:
            xx = self.transformer_encoder(xx, mask=mask)
        else:
            xx = self.transformer_encoder(xx)
        # return x + xx
        return x + xx.permute(1, 0, 2)  # (batch_size, seq_len, vocab_size)