ExpIvme-DiffusionConversate-v1 / modeling_expivme_diffusion.py
ereniko's picture
Add missing modeling_expivme_diffusion.py (auto_map target)
9b0fe46 verified
Raw
History Blame Contribute Delete
7.02 kB
"""HuggingFace Transformers model for ExpIvme-DiffusionConversate-v1.
A masked/absorbing-state discrete diffusion language model. Architecture
(RMSNorm, RoPE, SwiGLU, tied embeddings) inherited from
IvmeLabs/Ivme-Conversate-v2-Base, scaled to ~130M params, with bidirectional
(non-causal) attention for masked diffusion.
"""
from dataclasses import dataclass
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import PreTrainedModel
from transformers.modeling_outputs import ModelOutput
try:
from .configuration_expivme_diffusion import ExpIvmeDiffusionConfig
except ImportError:
from configuration_expivme_diffusion import ExpIvmeDiffusionConfig
def _precompute_rope_freqs(head_dim, max_seq_len, theta, device=None):
freqs = 1.0 / (theta ** (torch.arange(0, head_dim, 2, device=device).float() / head_dim))
positions = torch.arange(max_seq_len, device=device).float()
angles = torch.outer(positions, freqs)
return torch.cos(angles), torch.sin(angles)
def _apply_rope(x, rope_cos_sin):
cos, sin = rope_cos_sin
B, H, T, D = x.shape
x1 = x[..., 0::2]
x2 = x[..., 1::2]
cos = cos.view(1, 1, T, D // 2).to(x.dtype)
sin = sin.view(1, 1, T, D // 2).to(x.dtype)
out1 = x1 * cos - x2 * sin
out2 = x1 * sin + x2 * cos
out = torch.stack([out1, out2], dim=-1).reshape(B, H, T, D)
return out.type_as(x)
class ExpIvmeRMSNorm(nn.Module):
def __init__(self, dim, eps=1e-5):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x):
dtype = x.dtype
x = x.float()
rms = torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
return (x * rms).to(dtype) * self.weight
class ExpIvmeSelfAttention(nn.Module):
def __init__(self, hidden_dim, n_heads, dropout=0.0):
super().__init__()
self.n_heads = n_heads
self.head_dim = hidden_dim // n_heads
self.dropout = dropout
self.q_proj = nn.Linear(hidden_dim, hidden_dim, bias=False)
self.k_proj = nn.Linear(hidden_dim, hidden_dim, bias=False)
self.v_proj = nn.Linear(hidden_dim, hidden_dim, bias=False)
self.out_proj = nn.Linear(hidden_dim, hidden_dim, bias=False)
def forward(self, x, rope, attn_mask=None):
B, T, C = x.shape
q = self.q_proj(x).view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
k = self.k_proj(x).view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
v = self.v_proj(x).view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
q = _apply_rope(q, rope)
k = _apply_rope(k, rope)
out = F.scaled_dot_product_attention(
q, k, v, attn_mask=attn_mask, is_causal=False,
dropout_p=self.dropout if self.training else 0.0,
)
out = out.transpose(1, 2).contiguous().view(B, T, C)
return self.out_proj(out)
class ExpIvmeSwiGLU(nn.Module):
def __init__(self, hidden_dim, ffn_mult):
super().__init__()
inner_dim = int(hidden_dim * ffn_mult * 2 / 3)
inner_dim = ((inner_dim + 7) // 8) * 8
self.gate_proj = nn.Linear(hidden_dim, inner_dim, bias=False)
self.up_proj = nn.Linear(hidden_dim, inner_dim, bias=False)
self.down_proj = nn.Linear(inner_dim, hidden_dim, bias=False)
def forward(self, x):
return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))
class ExpIvmeBlock(nn.Module):
def __init__(self, hidden_dim, n_heads, ffn_mult, norm_eps, dropout=0.0):
super().__init__()
self.attn_norm = ExpIvmeRMSNorm(hidden_dim, eps=norm_eps)
self.attn = ExpIvmeSelfAttention(hidden_dim, n_heads, dropout)
self.ffn_norm = ExpIvmeRMSNorm(hidden_dim, eps=norm_eps)
self.ffn = ExpIvmeSwiGLU(hidden_dim, ffn_mult)
def forward(self, x, rope, attn_mask=None):
x = x + self.attn(self.attn_norm(x), rope, attn_mask=attn_mask)
x = x + self.ffn(self.ffn_norm(x))
return x
@dataclass
class DiffusionLMOutput(ModelOutput):
loss: torch.FloatTensor = None
logits: torch.FloatTensor = None
class ExpIvmeForDiffusionLMHub(PreTrainedModel):
"""Single module tree — self.model.* and self.lm_head only."""
config_class = ExpIvmeDiffusionConfig
base_model_prefix = "model"
_tied_weights_keys = {"lm_head.weight": "model.tok_embed.weight"}
def __init__(self, config):
super().__init__(config)
self.model = nn.Module()
self.model.tok_embed = nn.Embedding(config.vocab_size, config.hidden_dim)
self.model.blocks = nn.ModuleList([
ExpIvmeBlock(config.hidden_dim, config.n_heads, config.ffn_mult, config.norm_eps, config.dropout)
for _ in range(config.n_layers)
])
self.model.final_norm = ExpIvmeRMSNorm(config.hidden_dim, eps=config.norm_eps)
self.lm_head = nn.Linear(config.hidden_dim, config.vocab_size, bias=False)
self.head_dim = config.hidden_dim // config.n_heads
self.rope_theta = config.rope_theta
self.post_init()
if config.tie_word_embeddings:
self.tie_weights()
def get_input_embeddings(self):
return self.model.tok_embed
def set_input_embeddings(self, value):
self.model.tok_embed = value
def get_output_embeddings(self):
return self.lm_head
def forward(self, input_ids, attention_mask=None, labels=None, mask_positions=None, t=None, return_dict=True, **kw):
B, T = input_ids.shape
rope = _precompute_rope_freqs(self.head_dim, T, self.rope_theta, device=input_ids.device)
sdpa_mask = None
if attention_mask is not None:
sdpa_mask = torch.zeros(B, 1, 1, T, dtype=torch.float32, device=input_ids.device)
sdpa_mask.masked_fill_(attention_mask[:, None, None, :] == 0, float("-inf"))
sdpa_mask = sdpa_mask.to(dtype=self.model.tok_embed.weight.dtype)
x = self.model.tok_embed(input_ids)
for block in self.model.blocks:
x = block(x, rope, attn_mask=sdpa_mask)
x = self.model.final_norm(x)
logits = self.lm_head(x)
loss = None
if labels is not None and mask_positions is not None:
ce = F.cross_entropy(
logits.view(-1, self.config.vocab_size), labels.view(-1), ignore_index=-100, reduction="none",
).view(B, T)
ce = ce * mask_positions.float()
per_example_loss = ce.sum(dim=1)
if t is not None:
weight = 1.0 / t.clamp(min=1e-3)
per_example_loss = per_example_loss * weight
n_masked = mask_positions.float().sum(dim=1).clamp(min=1.0)
loss = (per_example_loss / n_masked).mean()
if not return_dict:
return (loss, logits) if loss is not None else (logits,)
return DiffusionLMOutput(loss=loss, logits=logits)
__all__ = ["ExpIvmeDiffusionConfig", "ExpIvmeForDiffusionLMHub"]