MCDP_GDN / mcdpgdn.py
Abdullah-Nazhat's picture
Update mcdpgdn.py
8d0c787 verified
Raw
History Blame Contribute Delete
4.73 kB
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)