Pebble-50M-beta / modeling_pebble.py
Hoglet-33's picture
Update modeling_pebble.py
56b4186 verified
Raw History Blame Contribute Delete
8.69 kB
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import PreTrainedModel, GenerationMixin
from transformers.modeling_outputs import CausalLMOutputWithPast
# Try to import the fast CUDA Mamba2.
try:
from mamba_ssm import Mamba2
HAS_MAMBA_SSM = True
except ImportError:
HAS_MAMBA_SSM = False
from .configuration_pebble import PebbleConfig
class RMSNorm(nn.Module):
def __init__(self, dim, eps=1e-6):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x):
dt = x.dtype
xf = x.float()
xf = xf * torch.rsqrt(xf.pow(2).mean(-1, keepdim=True) + self.eps)
return self.weight * xf.to(dt)
class AttentionBlock(nn.Module):
def __init__(self, config):
super().__init__()
dim = config.hidden_size
n_heads = config.num_attention_heads
hidden = config.intermediate_size
assert dim % n_heads == 0
self.nh, self.hd = n_heads, dim // n_heads
self.wqkv = nn.Linear(dim, 3 * dim, bias=False)
self.wo = nn.Linear(dim, dim, bias=False)
self.fc1 = nn.Linear(dim, hidden, bias=False)
self.fc2 = nn.Linear(hidden, dim, bias=False)
self.ln1 = RMSNorm(dim, eps=config.rms_norm_eps)
self.ln2 = RMSNorm(dim, eps=config.rms_norm_eps)
self.rope_theta = config.attention.get("rope_theta", 10000.0)
def forward(self, x):
B, T, C = x.shape
h = self.ln1(x)
qkv = self.wqkv(h).view(B, T, 3, self.nh, self.hd) \
.permute(2, 0, 3, 1, 4)
q, k, v = qkv[0], qkv[1], qkv[2]
half = self.hd // 2
invf = 1.0 / (self.rope_theta ** (
torch.arange(0, half, device=x.device, dtype=torch.float32)
* 2.0 / self.hd))
ang = torch.outer(
torch.arange(T, device=x.device, dtype=torch.float32), invf)
cos, sin = ang.cos()[None, None], ang.sin()[None, None]
q1, q2 = q.float()[..., :half], q.float()[..., half:]
k1, k2 = k.float()[..., :half], k.float()[..., half:]
q = torch.cat([q1 * cos - q2 * sin,
q1 * sin + q2 * cos], dim=-1).to(v.dtype)
k = torch.cat([k1 * cos - k2 * sin,
k1 * sin + k2 * cos], dim=-1).to(v.dtype)
y = F.scaled_dot_product_attention(q, k, v, is_causal=True)
y = y.transpose(1, 2).reshape(B, T, C)
x = x + self.wo(y)
x = x + self.fc2(F.gelu(self.fc1(self.ln2(x))))
return x
class PurePyTorchMamba2(nn.Module):
"""
A pure PyTorch implementation of Mamba2 that exactly matches the
parameter names and math of mamba_ssm.Mamba2, allowing it to run on CPU.
"""
def __init__(self, config):
super().__init__()
mamba_cfg = config.mamba2
d_model = config.hidden_size
d_state = mamba_cfg.get("d_state", 128)
d_conv = mamba_cfg.get("d_conv", 4)
expand = mamba_cfg.get("expand", 2)
headdim = mamba_cfg.get("headdim", 64)
self.d_model = d_model
self.d_state = d_state
self.d_inner = expand * d_model
self.headdim = headdim
self.nheads = self.d_inner // headdim
self.d_conv = d_conv
self.in_proj = nn.Linear(d_model, 2 * self.d_inner + 2 * d_state + self.nheads, bias=False)
self.out_proj = nn.Linear(self.d_inner, d_model, bias=False)
self.norm = RMSNorm(self.d_inner, eps=config.rms_norm_eps)
conv_in_channels = self.d_inner + 2 * d_state
self.conv1d = nn.Conv1d(
in_channels=conv_in_channels,
out_channels=conv_in_channels,
kernel_size=d_conv,
padding=d_conv-1,
groups=conv_in_channels,
bias=True
)
self.A_log = nn.Parameter(torch.zeros(self.nheads, dtype=torch.float32))
self.D = nn.Parameter(torch.ones(self.nheads))
self.dt_bias = nn.Parameter(torch.ones(self.nheads))
def forward(self, x):
B, T, C = x.shape
xzbc = self.in_proj(x)
x, z, B_p, C_p, dt = torch.split(
xzbc,
[self.d_inner, self.d_inner, self.d_state, self.d_state, self.nheads],
dim=-1
)
xbc = torch.cat([x, B_p, C_p], dim=-1)
xbc = xbc.transpose(1, 2)
xbc = self.conv1d(xbc)[:, :, :T]
xbc = xbc.transpose(1, 2)
x, B_p, C_p = torch.split(
xbc,
[self.d_inner, self.d_state, self.d_state],
dim=-1
)
x = F.silu(x)
x = self.norm(x)
z = self.norm(z)
dt = F.softplus(dt + self.dt_bias)
A = -torch.exp(self.A_log)
# Reshape x for multi-head SSM
x = x.view(B, T, self.nheads, self.headdim)
B_p = B_p.view(B, T, 1, 1, self.d_state)
C_p = C_p.view(B, T, 1, 1, self.d_state)
A = A.view(1, self.nheads, 1, 1)
h = torch.zeros(B, self.nheads, self.d_state, self.headdim, device=x.device)
ys = []
for t in range(T):
dt_t = dt[:, t].view(B, self.nheads, 1, 1)
dA = torch.exp(dt_t * A)
dB = dt_t * B_p[:, t]
x_t = x[:, t].unsqueeze(2)
h = dA * h + (dB.transpose(-1, -2) * x_t)
y = (h * C_p[:, t].transpose(-1, -2)).sum(dim=2)
ys.append(y)
y = torch.stack(ys, dim=1)
y = y.view(B, T, self.d_inner)
D = self.D.view(1, 1, self.nheads, 1)
y = y + (x * D).view(B, T, self.d_inner)
# z is already (B, T, d_inner), so it multiplies perfectly!
out = y * z
out = self.out_proj(out)
return out
class MambaBlock(nn.Module):
def __init__(self, config, layer_idx=0):
super().__init__()
self.ln = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
mamba_cfg = config.mamba2
if HAS_MAMBA_SSM and torch.cuda.is_available():
# Fast CUDA path for NVIDIA GPUs
self.use_hf_fallback = False
self.mixer = Mamba2(
d_model=config.hidden_size,
d_state=mamba_cfg.get("d_state", 128),
d_conv=mamba_cfg.get("d_conv", 4),
expand=mamba_cfg.get("expand", 2),
headdim=mamba_cfg.get("headdim", 64),
use_mem_eff_path=mamba_cfg.get("use_mem_eff_path", True),
)
else:
# Pure PyTorch fallback for Macs / CPUs / AMD GPUs
self.use_hf_fallback = True
self.mixer = PurePyTorchMamba2(config)
def forward(self, x):
return x + self.mixer(self.ln(x))
class PebbleForCausalLM(PreTrainedModel, GenerationMixin):
config_class = PebbleConfig
supports_gradient_checkpointing = False
_no_split_modules = ["MambaBlock", "AttentionBlock"]
def __init__(self, config):
super().__init__(config)
self.config = config
self.wte = nn.Embedding(config.vocab_size, config.hidden_size)
self.blocks = nn.ModuleList([
MambaBlock(config, layer_idx=i) if i % 4 < 3
else AttentionBlock(config)
for i in range(config.num_hidden_layers)
])
self.lnf = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
self.tie_weights()
def tie_weights(self):
if self.config.tie_word_embeddings:
self.lm_head.weight = self.wte.weight
def forward(self, input_ids=None, attention_mask=None, labels=None, past_key_values=None, **kwargs):
x = self.wte(input_ids)
for blk in self.blocks:
x = blk(x)
logits = self.lm_head(self.lnf(x))
loss = None
if labels is not None:
shift_logits = logits[..., :-1, :].contiguous()
shift_labels = labels[..., 1:].contiguous()
loss = F.cross_entropy(
shift_logits.view(-1, shift_logits.size(-1)),
shift_labels.view(-1)
)
return CausalLMOutputWithPast(
loss=loss,
logits=logits,
past_key_values=past_key_values,
)
def prepare_inputs_for_generation(self, input_ids, past_key_values=None, **kwargs):
return {
"input_ids": input_ids,
"past_key_values": past_key_values,
}