maba-v1.5-101m-test / maba_sparse /baselines /dense_transformer.py
AndrewThompson1233's picture
commit2
c54d56e
Raw
History Blame Contribute Delete
5.66 kB
import math
from typing import Any, List, Optional, Tuple, Union
import torch
import torch.nn as nn
import torch.nn.functional as F
from maba_sparse.model import FactorizedEmbeddings, MabaSparseOutput, RMSNorm, SwiGLUFFN
class DenseAttention(nn.Module):
def __init__(self, dim: int = 640, n_heads: int = 10, d_head: int = 64) -> None:
super().__init__()
self.dim = dim
self.n_heads = n_heads
self.d_head = d_head
self.scale = 1.0 / math.sqrt(d_head)
self.q_proj = nn.Linear(dim, n_heads * d_head, bias=False)
self.k_proj = nn.Linear(dim, n_heads * d_head, bias=False)
self.v_proj = nn.Linear(dim, n_heads * d_head, bias=False)
self.o_proj = nn.Linear(n_heads * d_head, dim, bias=False)
def forward(
self,
x: torch.Tensor,
kv_cache: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
) -> Tuple[torch.Tensor, Optional[Tuple[torch.Tensor, torch.Tensor]]]:
b, l, d = x.shape
q = self.q_proj(x).view(b, l, self.n_heads, self.d_head).transpose(1, 2)
k = self.k_proj(x).view(b, l, self.n_heads, self.d_head).transpose(1, 2)
v = self.v_proj(x).view(b, l, self.n_heads, self.d_head).transpose(1, 2)
if kv_cache is not None:
pk, pv = kv_cache
k = torch.cat([pk, k], dim=2)
v = torch.cat([pv, v], dim=2)
nkv = (k, v)
c = (kv_cache is None) and (l > 1)
o = F.scaled_dot_product_attention(q, k, v, is_causal=c)
o = o.transpose(1, 2).contiguous().view(b, l, self.n_heads * self.d_head)
return self.o_proj(o), nkv
class DenseTransformerBlock(nn.Module):
def __init__(
self,
dim: int = 640,
n_heads: int = 10,
d_head: int = 64,
intermediate_size: int = 1728,
eps: float = 1e-6,
residual_gate_bias: float = 2.0,
) -> None:
super().__init__()
self.norm1 = RMSNorm(dim, eps=eps)
self.mixer = DenseAttention(dim, n_heads, d_head)
self.res_gate1 = nn.Parameter(torch.full((dim,), residual_gate_bias))
self.norm2 = RMSNorm(dim, eps=eps)
self.ffn = SwiGLUFFN(dim, intermediate_size)
self.res_gate2 = nn.Parameter(torch.full((dim,), residual_gate_bias))
def forward(
self,
x: torch.Tensor,
kv_cache: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
) -> Tuple[torch.Tensor, Optional[Tuple[torch.Tensor, torch.Tensor]]]:
h = self.norm1(x)
ao, nkv = self.mixer(h, kv_cache=kv_cache)
x = x + torch.sigmoid(self.res_gate1) * ao
x = x + torch.sigmoid(self.res_gate2) * self.ffn(self.norm2(x))
return x, nkv
class DenseTransformerForCausalLM(nn.Module):
def __init__(
self,
vocab_size: int = 32768,
d_emb: int = 128,
dim: int = 640,
n_layers: int = 20,
n_heads: int = 10,
d_head: int = 64,
intermediate_size: int = 1728,
eps: float = 1e-6,
residual_gate_bias: float = 2.0,
) -> None:
super().__init__()
self.vocab_size = vocab_size
self.dim = dim
self.n_layers = n_layers
self.embeddings = FactorizedEmbeddings(vocab_size, d_emb, dim)
self.layers = nn.ModuleList([
DenseTransformerBlock(
dim=dim,
n_heads=n_heads,
d_head=d_head,
intermediate_size=intermediate_size,
eps=eps,
residual_gate_bias=residual_gate_bias,
)
for _ in range(n_layers)
])
self.final_norm = RMSNorm(dim, eps=eps)
self.head_proj = nn.Linear(dim, d_emb, bias=False)
self.lm_head = nn.Linear(d_emb, vocab_size, bias=False)
self.lm_head.weight = self.embeddings.in_emb.weight
def forward(
self,
input_ids: torch.Tensor,
targets: Optional[torch.Tensor] = None,
labels: Optional[torch.Tensor] = None,
past_states: Optional[List[Any]] = None,
) -> MabaSparseOutput:
if targets is None and labels is not None:
targets = labels
x = self.embeddings(input_ids)
nps = []
for i, layer in enumerate(self.layers):
kv = past_states[i] if past_states is not None else None
x, nkv = layer(x, kv_cache=kv)
nps.append(nkv)
xn = self.final_norm(x)
logits = self.lm_head(self.head_proj(xn))
loss = None
if targets is not None:
loss = F.cross_entropy(logits.view(-1, self.vocab_size), targets.view(-1))
return MabaSparseOutput(
logits=logits,
loss=loss,
past_states=nps,
)
@torch.no_grad()
def generate(
self,
input_ids: torch.Tensor,
max_new_tokens: int = 32,
temperature: float = 1.0,
top_k: Optional[int] = 50,
) -> torch.Tensor:
self.eval()
gen = input_ids.clone()
for _ in range(max_new_tokens):
out = self(gen)
nl = out.logits[:, -1, :]
if temperature > 0:
nl = nl / temperature
if top_k is not None:
v, _ = torch.topk(nl, min(top_k, nl.size(-1)))
nl[nl < v[:, [-1]]] = float("-inf")
p = F.softmax(nl, dim=-1)
tok = torch.multinomial(p, num_samples=1)
else:
tok = torch.argmax(nl, dim=-1, keepdim=True)
gen = torch.cat([gen, tok], dim=1)
return gen
DenseTransformerLM = DenseTransformerForCausalLM