Abdullah-Nazhat commited on
Commit
8d0c787
·
verified ·
1 Parent(s): 7769ef0

Update mcdpgdn.py

Browse files
Files changed (1) hide show
  1. 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
- # We could implement this as
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) # NEW
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)