| import torch |
| from torch import nn |
|
|
| def l2norm(x, dim=-1, eps=1e-6): |
| return x * torch.rsqrt((x * x).sum(dim=dim, keepdim=True) + eps) |
|
|
| class GatedDeltaNet(nn.Module): |
| def __init__( |
| self, d_in, d_out, dropout, num_heads, qkv_bias=False |
| ): |
| super().__init__() |
| assert d_out % num_heads == 0 |
|
|
| self.d_out = d_out |
| self.num_heads = num_heads |
| self.head_dim = d_out // num_heads |
|
|
| self.W_query = nn.Linear(d_in, d_out, bias=qkv_bias) |
| self.W_key = nn.Linear(d_in, d_out, bias=qkv_bias) |
| self.W_value = nn.Linear(d_in, d_out, bias=qkv_bias) |
| |
| self.W_gate = nn.Linear(d_in, d_out, bias=False) |
| self.W_beta = nn.Linear(d_in, d_out, bias=False) |
|
|
| self.W_alpha = nn.Linear(d_in, num_heads, bias=False) |
| self.dt_bias = nn.Parameter(torch.ones(num_heads)) |
| A_init = torch.empty(num_heads).uniform_(0, 16) |
| self.A_log = nn.Parameter(torch.log(A_init)) |
| |
| self.norm = nn.RMSNorm(self.head_dim, eps=1e-6) |
|
|
| self.out_proj = nn.Linear(d_out, d_out) |
| self.dropout = nn.Dropout(dropout) |
|
|
| def forward(self, x): |
| b, num_tokens, _ = x.shape |
| queries = self.W_query(x) |
| keys = self.W_key(x) |
| values = self.W_value(x) |
| |
| beta = torch.sigmoid(self.W_beta(x)) |
| alpha_log = -self.A_log.exp().view(1, 1, -1) * F.softplus( |
| self.W_alpha(x) + self.dt_bias |
| ) |
| alpha = alpha_log.exp() |
| gate = self.W_gate(x) |
| |
| keys = keys.view(b, num_tokens, self.num_heads, self.head_dim) |
| values = values.view(b, num_tokens, self.num_heads, self.head_dim) |
| queries = queries.view(b, num_tokens, self.num_heads, self.head_dim) |
| beta = beta.view(b, num_tokens, self.num_heads, self.head_dim) |
| gate = gate.view(b, num_tokens, self.num_heads, self.head_dim) |
|
|
| keys = keys.transpose(1, 2) |
| queries = queries.transpose(1, 2) |
| values = values.transpose(1, 2) |
| beta = beta.transpose(1, 2) |
|
|
| queries = l2norm(queries, dim=-1) / (self.head_dim ** 0.5) |
| keys = l2norm(keys, dim=-1) |
| |
| S = x.new_zeros(b, self.num_heads, self.head_dim, self.head_dim) |
|
|
| outs = [] |
| |
| for t in range(num_tokens): |
| k_t = keys[:, :, t] |
| q_t = queries[:, :, t] |
| v_t = values[:, :, t] |
| b_t = beta[:, :, t] |
| a_t = alpha[:, t].unsqueeze(-1).unsqueeze(-1) |
|
|
| S = S * a_t |
| kv_mem = (S * k_t.unsqueeze(-1)).sum(dim=-2) |
| delta = (v_t - kv_mem) * b_t |
| S = S + k_t.unsqueeze(-1) * delta.unsqueeze(-2) |
| y_t = (S * q_t.unsqueeze(-1)).sum(dim=-2) |
| |
| outs.append(y_t) |
|
|
| context = torch.stack(outs, dim=2).transpose(1, 2).contiguous() |
| context = context.view(b, num_tokens, self.num_heads, self.head_dim) |
|
|
| context = self.norm(context) |
| context = context * F.silu(gate) |
| |
| context = context.view(b, num_tokens, self.d_out) |
| context = self.dropout(context) |
| out = self.out_proj(context) |
| return out |
| |
| class FeedForward(nn.Module): |
| def __init__(self, dim, hidden_dim, dropout): |
| super().__init__() |
| self.net = nn.Sequential( |
| nn.Linear(dim, hidden_dim), |
| nn.GELU(), |
| nn.Dropout(dropout), |
| nn.Linear(hidden_dim, dim), |
| nn.Dropout(dropout) |
| ) |
| def forward(self, x): |
| return self.net(x) |
|
|
| class MCGatingUnit(nn.Module): |
| def __init__(self,dim,dropout): |
| super().__init__() |
| |
| self.gdn_1 = GatedDeltaNet(dim,dim,dropout,8) |
| self.gdn_2 = GatedDeltaNet(dim,dim,dropout,8) |
| |
| def forward(self, x): |
| u, v = x, x |
| u = self.gdn_1(u) |
| v = self.gdn_1(v) |
| out = u * v |
| return out |
|
|
| class MCDPGDNBlock(nn.Module): |
| def __init__(self, d_model, d_ffn, dropout): |
| super().__init__() |
| |
| self.norm = nn.LayerNorm(d_model) |
| self.mcgu = MCGatingUnit(d_model,dropout) |
| self.ffn = FeedForward(d_model,d_ffn,dropout) |
| |
| def forward(self, x): |
| residual = x |
| x = self.norm(x) |
| x = self.mcgu(x) |
| x = x + residual |
| residual = x |
| x = self.norm(x) |
| x = self.ffn(x) |
| out = x + residual |
| return out |
|
|
| class MCDPGDN(nn.Module): |
| def __init__(self, d_model, d_ffn, num_layers, dropout): |
| super().__init__() |
| |
| self.model = nn.Sequential( |
| *[MCDPGDNBlock(d_model, d_ffn, dropout) for _ in range(num_layers)] |
| ) |
|
|
| def forward(self, x): |
| return self.model(x) |