Update mcdpgdn.py
Browse files- mcdpgdn.py +9 -32
mcdpgdn.py
CHANGED
|
@@ -1,11 +1,5 @@
|
|
| 1 |
import torch
|
| 2 |
-
import numpy as np
|
| 3 |
from torch import nn
|
| 4 |
-
from torch.nn import functional as F
|
| 5 |
-
import einops
|
| 6 |
-
from einops.layers.torch import Rearrange
|
| 7 |
-
import math
|
| 8 |
-
|
| 9 |
|
| 10 |
def l2norm(x, dim=-1, eps=1e-6):
|
| 11 |
return x * torch.rsqrt((x * x).sum(dim=dim, keepdim=True) + eps)
|
|
@@ -24,24 +18,16 @@ class GatedDeltaNet(nn.Module):
|
|
| 24 |
self.W_query = nn.Linear(d_in, d_out, bias=qkv_bias)
|
| 25 |
self.W_key = nn.Linear(d_in, d_out, bias=qkv_bias)
|
| 26 |
self.W_value = nn.Linear(d_in, d_out, bias=qkv_bias)
|
| 27 |
-
|
| 28 |
-
### NEW: Gates for delta rule and output gating
|
| 29 |
self.W_gate = nn.Linear(d_in, d_out, bias=False)
|
| 30 |
self.W_beta = nn.Linear(d_in, d_out, bias=False)
|
| 31 |
|
| 32 |
-
# Note: The decay gate alpha corresponds to
|
| 33 |
-
# A_log + W_alpha(x) + dt_bias
|
| 34 |
self.W_alpha = nn.Linear(d_in, num_heads, bias=False)
|
| 35 |
self.dt_bias = nn.Parameter(torch.ones(num_heads))
|
| 36 |
A_init = torch.empty(num_heads).uniform_(0, 16)
|
| 37 |
self.A_log = nn.Parameter(torch.log(A_init))
|
| 38 |
-
|
| 39 |
-
# W_alpha = nn.Linear(d_in, num_heads, bias=True)
|
| 40 |
-
# but the bias is separate for interpretability and
|
| 41 |
-
# to mimic the official implementation
|
| 42 |
-
|
| 43 |
self.norm = nn.RMSNorm(self.head_dim, eps=1e-6)
|
| 44 |
-
####################################################
|
| 45 |
|
| 46 |
self.out_proj = nn.Linear(d_out, d_out)
|
| 47 |
self.dropout = nn.Dropout(dropout)
|
|
@@ -51,38 +37,32 @@ class GatedDeltaNet(nn.Module):
|
|
| 51 |
queries = self.W_query(x)
|
| 52 |
keys = self.W_key(x)
|
| 53 |
values = self.W_value(x)
|
| 54 |
-
|
| 55 |
-
### NEW: Compute delta rule gates
|
| 56 |
beta = torch.sigmoid(self.W_beta(x))
|
| 57 |
alpha_log = -self.A_log.exp().view(1, 1, -1) * F.softplus(
|
| 58 |
self.W_alpha(x) + self.dt_bias
|
| 59 |
)
|
| 60 |
alpha = alpha_log.exp()
|
| 61 |
gate = self.W_gate(x)
|
| 62 |
-
|
| 63 |
-
|
| 64 |
keys = keys.view(b, num_tokens, self.num_heads, self.head_dim)
|
| 65 |
values = values.view(b, num_tokens, self.num_heads, self.head_dim)
|
| 66 |
queries = queries.view(b, num_tokens, self.num_heads, self.head_dim)
|
| 67 |
beta = beta.view(b, num_tokens, self.num_heads, self.head_dim)
|
| 68 |
-
gate = gate.view(b, num_tokens, self.num_heads, self.head_dim)
|
| 69 |
|
| 70 |
keys = keys.transpose(1, 2)
|
| 71 |
queries = queries.transpose(1, 2)
|
| 72 |
values = values.transpose(1, 2)
|
| 73 |
beta = beta.transpose(1, 2)
|
| 74 |
|
| 75 |
-
####################################################
|
| 76 |
-
### NEW: QKNorm-like normalization for delta rule
|
| 77 |
queries = l2norm(queries, dim=-1) / (self.head_dim ** 0.5)
|
| 78 |
keys = l2norm(keys, dim=-1)
|
| 79 |
-
|
| 80 |
-
|
| 81 |
S = x.new_zeros(b, self.num_heads, self.head_dim, self.head_dim)
|
| 82 |
|
| 83 |
outs = []
|
| 84 |
-
|
| 85 |
-
### NEW: Gated delta rule update
|
| 86 |
for t in range(num_tokens):
|
| 87 |
k_t = keys[:, :, t]
|
| 88 |
q_t = queries[:, :, t]
|
|
@@ -95,18 +75,15 @@ class GatedDeltaNet(nn.Module):
|
|
| 95 |
delta = (v_t - kv_mem) * b_t
|
| 96 |
S = S + k_t.unsqueeze(-1) * delta.unsqueeze(-2)
|
| 97 |
y_t = (S * q_t.unsqueeze(-1)).sum(dim=-2)
|
| 98 |
-
|
| 99 |
outs.append(y_t)
|
| 100 |
|
| 101 |
context = torch.stack(outs, dim=2).transpose(1, 2).contiguous()
|
| 102 |
context = context.view(b, num_tokens, self.num_heads, self.head_dim)
|
| 103 |
|
| 104 |
-
####################################################
|
| 105 |
-
### NEW: Apply RMSNorm and SiLU gate
|
| 106 |
context = self.norm(context)
|
| 107 |
context = context * F.silu(gate)
|
| 108 |
-
|
| 109 |
-
|
| 110 |
context = context.view(b, num_tokens, self.d_out)
|
| 111 |
context = self.dropout(context)
|
| 112 |
out = self.out_proj(context)
|
|
|
|
| 1 |
import torch
|
|
|
|
| 2 |
from torch import nn
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 3 |
|
| 4 |
def l2norm(x, dim=-1, eps=1e-6):
|
| 5 |
return x * torch.rsqrt((x * x).sum(dim=dim, keepdim=True) + eps)
|
|
|
|
| 18 |
self.W_query = nn.Linear(d_in, d_out, bias=qkv_bias)
|
| 19 |
self.W_key = nn.Linear(d_in, d_out, bias=qkv_bias)
|
| 20 |
self.W_value = nn.Linear(d_in, d_out, bias=qkv_bias)
|
| 21 |
+
|
|
|
|
| 22 |
self.W_gate = nn.Linear(d_in, d_out, bias=False)
|
| 23 |
self.W_beta = nn.Linear(d_in, d_out, bias=False)
|
| 24 |
|
|
|
|
|
|
|
| 25 |
self.W_alpha = nn.Linear(d_in, num_heads, bias=False)
|
| 26 |
self.dt_bias = nn.Parameter(torch.ones(num_heads))
|
| 27 |
A_init = torch.empty(num_heads).uniform_(0, 16)
|
| 28 |
self.A_log = nn.Parameter(torch.log(A_init))
|
| 29 |
+
|
|
|
|
|
|
|
|
|
|
|
|
|
| 30 |
self.norm = nn.RMSNorm(self.head_dim, eps=1e-6)
|
|
|
|
| 31 |
|
| 32 |
self.out_proj = nn.Linear(d_out, d_out)
|
| 33 |
self.dropout = nn.Dropout(dropout)
|
|
|
|
| 37 |
queries = self.W_query(x)
|
| 38 |
keys = self.W_key(x)
|
| 39 |
values = self.W_value(x)
|
| 40 |
+
|
|
|
|
| 41 |
beta = torch.sigmoid(self.W_beta(x))
|
| 42 |
alpha_log = -self.A_log.exp().view(1, 1, -1) * F.softplus(
|
| 43 |
self.W_alpha(x) + self.dt_bias
|
| 44 |
)
|
| 45 |
alpha = alpha_log.exp()
|
| 46 |
gate = self.W_gate(x)
|
| 47 |
+
|
|
|
|
| 48 |
keys = keys.view(b, num_tokens, self.num_heads, self.head_dim)
|
| 49 |
values = values.view(b, num_tokens, self.num_heads, self.head_dim)
|
| 50 |
queries = queries.view(b, num_tokens, self.num_heads, self.head_dim)
|
| 51 |
beta = beta.view(b, num_tokens, self.num_heads, self.head_dim)
|
| 52 |
+
gate = gate.view(b, num_tokens, self.num_heads, self.head_dim)
|
| 53 |
|
| 54 |
keys = keys.transpose(1, 2)
|
| 55 |
queries = queries.transpose(1, 2)
|
| 56 |
values = values.transpose(1, 2)
|
| 57 |
beta = beta.transpose(1, 2)
|
| 58 |
|
|
|
|
|
|
|
| 59 |
queries = l2norm(queries, dim=-1) / (self.head_dim ** 0.5)
|
| 60 |
keys = l2norm(keys, dim=-1)
|
| 61 |
+
|
|
|
|
| 62 |
S = x.new_zeros(b, self.num_heads, self.head_dim, self.head_dim)
|
| 63 |
|
| 64 |
outs = []
|
| 65 |
+
|
|
|
|
| 66 |
for t in range(num_tokens):
|
| 67 |
k_t = keys[:, :, t]
|
| 68 |
q_t = queries[:, :, t]
|
|
|
|
| 75 |
delta = (v_t - kv_mem) * b_t
|
| 76 |
S = S + k_t.unsqueeze(-1) * delta.unsqueeze(-2)
|
| 77 |
y_t = (S * q_t.unsqueeze(-1)).sum(dim=-2)
|
| 78 |
+
|
| 79 |
outs.append(y_t)
|
| 80 |
|
| 81 |
context = torch.stack(outs, dim=2).transpose(1, 2).contiguous()
|
| 82 |
context = context.view(b, num_tokens, self.num_heads, self.head_dim)
|
| 83 |
|
|
|
|
|
|
|
| 84 |
context = self.norm(context)
|
| 85 |
context = context * F.silu(gate)
|
| 86 |
+
|
|
|
|
| 87 |
context = context.view(b, num_tokens, self.d_out)
|
| 88 |
context = self.dropout(context)
|
| 89 |
out = self.out_proj(context)
|