MetaDiffusion-150M-exp / inference.py
CodeSoft's picture
Upload 6 files
249471c verified
Raw
History Blame Contribute Delete
14.8 kB
#!/usr/bin/env python3
import argparse
import json
import os
import sys
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from safetensors.torch import load_file
from transformers import AutoTokenizer, AutoConfig
# ---------------------------------------------------------------------------
# Model definition
# ---------------------------------------------------------------------------
class MetaDiffusionConfig:
def __init__(self, **kwargs):
for k, v in kwargs.items():
setattr(self, k, v)
class RMSNorm(nn.Module):
def __init__(self, hidden_size, eps=1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(hidden_size))
self.eps = eps
def forward(self, x):
var = x.pow(2).mean(-1, keepdim=True)
x = x * torch.rsqrt(var + self.eps)
return self.weight * x
class RotaryEmbedding(nn.Module):
def __init__(self, dim, max_position_embeddings=5120, base=10000.0):
super().__init__()
inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2).float() / dim))
self.register_buffer("inv_freq", inv_freq, persistent=False)
def forward(self, x, position_ids):
inv_freq_expanded = self.inv_freq[None, :, None].float().expand(
position_ids.shape[0], -1, 1
)
position_ids_expanded = position_ids[:, None, :].float()
freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(1, 2)
emb = torch.cat((freqs, freqs), dim=-1)
return emb.cos().to(dtype=x.dtype), emb.sin().to(dtype=x.dtype)
def rotate_half(x):
x1, x2 = x.chunk(2, dim=-1)
return torch.cat((-x2, x1), dim=-1)
def apply_rotary_pos_emb(q, k, cos, sin):
cos = cos.unsqueeze(1)
sin = sin.unsqueeze(1)
q_embed = (q * cos) + (rotate_half(q) * sin)
k_embed = (k * cos) + (rotate_half(k) * sin)
return q_embed, k_embed
class TimestepEmbedding(nn.Module):
def __init__(self, hidden_size):
super().__init__()
self.mlp = nn.Sequential(
nn.Linear(hidden_size, hidden_size * 4),
nn.SiLU(),
nn.Linear(hidden_size * 4, hidden_size),
)
def forward(self, t):
half_dim = self.mlp[0].in_features // 2
emb = math.log(10000.0) / (half_dim - 1)
emb = torch.exp(torch.arange(half_dim, device=t.device) * -emb)
emb = t[:, None].float() * emb[None, :]
emb = torch.cat([emb.sin(), emb.cos()], dim=-1)
return self.mlp(emb)
class TimestepResidual(nn.Module):
def __init__(self, hidden_size):
super().__init__()
self.proj = nn.Linear(hidden_size, hidden_size)
def forward(self, x, emb):
return x + self.proj(emb)[:, None, :]
class SelfAttention(nn.Module):
def __init__(self, config):
super().__init__()
self.hidden_size = config.hidden_size
self.num_heads = config.num_attention_heads
self.num_kv_heads = config.num_key_value_heads
self.head_dim = config.head_dim
self.num_kv_groups = self.num_heads // self.num_kv_heads
self.q_proj = nn.Linear(config.hidden_size, self.num_heads * self.head_dim, bias=False)
self.k_proj = nn.Linear(config.hidden_size, self.num_kv_heads * self.head_dim, bias=False)
self.v_proj = nn.Linear(config.hidden_size, self.num_kv_heads * self.head_dim, bias=False)
self.o_proj = nn.Linear(self.num_heads * self.head_dim, config.hidden_size, bias=False)
self.rotary_emb = RotaryEmbedding(
config.head_dim,
max_position_embeddings=config.max_position_embeddings,
base=config.rope_theta,
)
def forward(self, x, position_ids):
batch, seq, _ = x.shape
q = self.q_proj(x).view(batch, seq, self.num_heads, self.head_dim).transpose(1, 2)
k = self.k_proj(x).view(batch, seq, self.num_kv_heads, self.head_dim).transpose(1, 2)
v = self.v_proj(x).view(batch, seq, self.num_kv_heads, self.head_dim).transpose(1, 2)
cos, sin = self.rotary_emb(x, position_ids)
q, k = apply_rotary_pos_emb(q, k, cos, sin)
if self.num_kv_groups > 1:
k = k.repeat_interleave(self.num_kv_groups, dim=1)
v = v.repeat_interleave(self.num_kv_groups, dim=1)
out = F.scaled_dot_product_attention(q, k, v)
out = out.transpose(1, 2).contiguous().view(batch, seq, -1)
return self.o_proj(out)
class MLP(nn.Module):
def __init__(self, config):
super().__init__()
self.gate_proj = nn.Linear(config.hidden_size, config.intermediate_size, bias=False)
self.up_proj = nn.Linear(config.hidden_size, config.intermediate_size, bias=False)
self.down_proj = nn.Linear(config.intermediate_size, config.hidden_size, bias=False)
def forward(self, x):
return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))
class TransformerBlock(nn.Module):
def __init__(self, config):
super().__init__()
self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.self_attn = SelfAttention(config)
self.post_attention_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.mlp = MLP(config)
self.timestep_residual = TimestepResidual(config.hidden_size)
def forward(self, x, timestep_emb, position_ids):
residual = x
x = self.input_layernorm(x)
x = self.self_attn(x, position_ids)
x = residual + x
x = self.timestep_residual(x, timestep_emb)
residual = x
x = self.post_attention_layernorm(x)
x = self.mlp(x)
x = residual + x
x = self.timestep_residual(x, timestep_emb)
return x
class MetaDiffusionLM(nn.Module):
def __init__(self, config):
super().__init__()
self.config = config
self.mask_token_id = getattr(config, "mask_token_id", config.vocab_size)
self.embed_tokens = nn.Embedding(
config.mask_vocab_size, config.hidden_size,
padding_idx=getattr(config, "pad_token_id", 1)
)
self.timestep_emb = TimestepEmbedding(getattr(config, "timestep_emb_hidden", config.hidden_size))
self.layers = nn.ModuleList([TransformerBlock(config) for _ in range(config.num_hidden_layers)])
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.lm_head = nn.Linear(config.hidden_size, config.mask_vocab_size, bias=False)
def forward(self, input_ids, timesteps):
batch, seq = input_ids.shape
position_ids = torch.arange(seq, device=input_ids.device).unsqueeze(0).expand(batch, -1)
x = self.embed_tokens(input_ids)
t_emb = self.timestep_emb(timesteps)
for layer in self.layers:
x = layer(x, t_emb, position_ids)
x = self.norm(x)
logits = self.lm_head(x)
return logits
# ---------------------------------------------------------------------------
# Generation
# ---------------------------------------------------------------------------
def cumulative_unmask_frac(i, N, schedule="cosine"):
if schedule == "cosine":
return 0.5 * (1 - math.cos(math.pi * i / N))
return i / N
def generate(model, tokenizer, prompt, seq_len=256, num_steps=64, device="cuda",
temperature=0.6, repetition_penalty=1.5, watch=False,
watch_every=1, mask_token_id=32000):
model.eval()
# Tokenize prompt
prompt_ids = tokenizer.encode(prompt, add_special_tokens=False)
prompt_ids = torch.tensor([prompt_ids], device=device)
# Build input: prompt + [MASK] tokens
total_len = seq_len
gen_len = max(total_len - prompt_ids.shape[1], 0)
input_ids = torch.full((1, total_len), mask_token_id, device=device, dtype=torch.long)
input_ids[0, :prompt_ids.shape[1]] = prompt_ids
for i in range(num_steps):
frac_now = cumulative_unmask_frac(i, num_steps)
frac_next = cumulative_unmask_frac(i + 1, num_steps)
# How many tokens to unmask this step
n_masked = (input_ids == mask_token_id).sum().item()
n_total_to_unmask = int((frac_next - frac_now) * (total_len - prompt_ids.shape[1]) + 0.5)
if i == num_steps - 1:
n_unmask = n_masked
else:
n_unmask = max(n_total_to_unmask, 1) if n_masked > 0 else 0
t = 1.0 - frac_now
t_batch = torch.full((1,), t, device=device)
with torch.no_grad():
logits = model(input_ids, t_batch)
# Prevent model from predicting [MASK] token
logits[:, :, mask_token_id] = -1e9
if repetition_penalty != 1.0:
for tok in input_ids[0].unique():
tok_idx = tok.item()
logits[0, :, tok_idx] = torch.where(
logits[0, :, tok_idx] < 0,
logits[0, :, tok_idx] * repetition_penalty,
logits[0, :, tok_idx] / repetition_penalty
)
# Sample at masked positions
mask_positions = (input_ids == mask_token_id)
mask_logits = logits[mask_positions]
probs = F.softmax(mask_logits / temperature, dim=-1)
sampled = torch.multinomial(probs, 1).squeeze(-1)
# Select which masks to fill (by confidence)
if n_unmask < mask_positions.sum():
# Get entropy/confidence for each mask
log_probs = F.log_softmax(mask_logits, dim=-1)
confidence, _ = log_probs.max(dim=-1)
_, top_indices = confidence.topk(n_unmask)
# Only fill top-confidence positions
mask_flat = mask_positions.nonzero(as_tuple=False)
fill_positions = mask_flat[top_indices]
for idx, tok in zip(fill_positions, sampled[top_indices]):
input_ids[idx[0], idx[1]] = tok
else:
# Fill all remaining masks
input_ids[mask_positions] = sampled
if watch and i % watch_every == 0:
text = tokenizer.decode(input_ids[0], skip_special_tokens=True)
n_remaining = (input_ids == mask_token_id).sum().item()
print(f"Step {i+1}/{num_steps} | LR={t:.3f} | Masks remaining: {n_remaining}")
print(text[:200])
print()
# Decode
return tokenizer.decode(input_ids[0], skip_special_tokens=False)
def load_model(model_path, device="cuda"):
"""Load model from safetensors file, directory, or HuggingFace Hub."""
# Check if it's a local path or HF hub id
is_file = os.path.isfile(model_path) and model_path.endswith(".safetensors")
is_dir = os.path.isdir(model_path)
is_local = is_file or is_dir
if is_local:
if is_file:
safetensors_path = model_path
config_path = os.path.join(os.path.dirname(model_path), "config.json")
else:
config_path = os.path.join(model_path, "config.json")
safetensors_path = os.path.join(model_path, "model.safetensors")
if not os.path.isfile(safetensors_path):
print(f"ERROR: model.safetensors not found in {model_path}")
sys.exit(1)
if not os.path.isfile(config_path):
print(f"ERROR: config.json not found next to {safetensors_path}")
sys.exit(1)
with open(config_path) as f:
config_dict = json.load(f)
else:
# Load from HuggingFace Hub
from huggingface_hub import hf_hub_download
config_path = hf_hub_download(model_path, "config.json")
safetensors_path = hf_hub_download(model_path, "model.safetensors")
with open(config_path) as f:
config_dict = json.load(f)
# Build config
config = MetaDiffusionConfig(**config_dict)
model = MetaDiffusionLM(config)
model = model.to(device)
# Load weights (remap HF names to model names)
state_dict = load_file(safetensors_path)
# Remap from HF naming to model naming
new_state_dict = {}
for key, value in state_dict.items():
if key.startswith("model."):
new_key = key[len("model."):]
else:
new_key = key
new_state_dict[new_key] = value
result = model.load_state_dict(new_state_dict, strict=False)
if result.missing_keys:
print(f" Warning: missing keys: {result.missing_keys[:5]}...")
if result.unexpected_keys:
print(f" Warning: unexpected keys: {result.unexpected_keys[:5]}...")
model = model.to(device)
print(f" Model loaded: {sum(p.numel() for p in model.parameters())/1e6:.1f}M params")
return model, config
def main():
parser = argparse.ArgumentParser(description="MetaDiffusion inference")
parser.add_argument("--model-path", required=True, help="Path to model directory or HF Hub ID")
parser.add_argument("--prompt", default="The cat sat on the", help="Input prompt")
parser.add_argument("--seq-len", type=int, default=256, help="Sequence length")
parser.add_argument("--num-steps", type=int, default=512, help="Denoising steps")
parser.add_argument("--temperature", type=float, default=0.6, help="Sampling temperature")
parser.add_argument("--repetition-penalty", type=float, default=1.5, help="Repetition penalty")
parser.add_argument("--device", default="cuda", help="Device (cuda/cpu)")
parser.add_argument("--watch", action="store_true", help="Show denoising progress")
parser.add_argument("--watch-every", type=int, default=4, help="Show progress every N steps")
parser.add_argument("--base-model", default="SupraLabs/Supra-1.5-50M-Base-exp",
help="HuggingFace model for tokenizer")
args = parser.parse_args()
if "cpu" in args.device:
device = torch.device("cpu")
else:
device = torch.device(args.device if torch.cuda.is_available() else "cpu")
model, config = load_model(args.model_path, device)
# Load tokenizer from base model
tokenizer = AutoTokenizer.from_pretrained(args.base_model)
mask_token_id = getattr(config, "mask_token_id", config.vocab_size)
print(f"\nPrompt: '{args.prompt}'")
print(f"Steps: {args.num_steps} | Temp: {args.temperature}")
print()
output = generate(
model, tokenizer, args.prompt,
seq_len=args.seq_len, num_steps=args.num_steps,
device=device, temperature=args.temperature,
repetition_penalty=args.repetition_penalty,
watch=args.watch,
watch_every=args.watch_every, mask_token_id=mask_token_id
)
print("Output:")
print(output)
if __name__ == "__main__":
main()