Voltline's picture
Release VimeML V2.1 step40000 FP32 and Core ML INT8 (GPL-2.0)
29f25be verified
Raw History Blame Contribute Delete
5.7 kB
"""Optional training compilation and frozen-weight initialization for V2."""
from types import MethodType
from pathlib import Path
import torch
from torch.nn import functional as F
from vimeml.training.model_v2 import RMSNorm
def rmsnorm_operations(hidden, weight, eps):
normalized = hidden.float()
variance = normalized.pow(2).mean(-1, keepdim=True)
normalized = normalized * torch.rsqrt(variance + eps)
return normalized.to(hidden.dtype) * weight
def fuse_rmsnorm(model):
"""Compile only pure norm operations, leaving variable-length logits eager."""
compiled = torch.compile(rmsnorm_operations, fullgraph=True, dynamic=True)
def forward(module, hidden):
return compiled(hidden, module.weight, module.eps)
count = 0
for module in model.modules():
if isinstance(module, RMSNorm):
module.forward = MethodType(forward, module)
count += 1
if not count:
raise ValueError("RMSNorm fusion requires a V2 model.")
return {
"implementation": "torch.compile RMSNorm only",
"norm_modules": count,
"dynamic_shapes": True,
"state_dict_unchanged": True,
"checkpoint_inference": "Plain primitive RMSNorm; training fusion is optional.",
}
def compile_backbone(model):
"""Compile the full hidden-state stack, keeping variable-size token loss eager."""
def hidden_stack(input_ids):
positions = torch.arange(input_ids.shape[1], device=input_ids.device)
hidden = model.dropout(
model.token_embedding(input_ids) + model.position_embedding(positions)
)
for block in model.blocks:
hidden = block(hidden)
return model.final_norm(hidden)
compiled = torch.compile(hidden_stack, fullgraph=True, dynamic=True)
def forward(module, input_ids, labels=None):
if (
input_ids.ndim != 2
or not 1 <= input_ids.shape[1] <= module.config.context_length
):
raise ValueError("Expected [batch, time] within the model context.")
hidden = compiled(input_ids)
if labels is None:
return module.lm_head(hidden)
if labels.shape != input_ids.shape:
raise ValueError("Labels must match input shape.")
valid = labels != -100
logits = module.lm_head(hidden[valid])
return {
"loss_sum": F.cross_entropy(logits, labels[valid], reduction="sum"),
"token_count": valid.sum(),
}
model.forward = MethodType(forward, model)
return {
"implementation": "torch.compile full hidden stack; eager valid-token head and CE",
"dynamic_shapes": True,
"state_dict_unchanged": True,
"checkpoint_inference": "Plain TinyGPTV2",
}
def initialize_weights(model, checkpoint, signatures, optimizer=None, precision=None):
"""Start a separate run; optionally carry compatible AdamW moments forward."""
path = Path(checkpoint)
saved = torch.load(path, map_location="cpu", weights_only=True)
if (
saved.get("format") != "vimeml_tiny_gpt_v2"
or saved["model_config"] != model.configuration()
):
raise ValueError("Initialization requires the same V2 model configuration.")
if any(
saved["signatures"][name] != signatures[name] for name in ("tokens", "windows")
):
raise ValueError(
"Initialization requires the same tokenizer/token store and window index."
)
if optimizer is not None:
if saved.get("precision") != precision or not saved.get("optimizer", {}).get(
"state"
):
raise ValueError(
"Optimizer continuation requires matching precision and nonempty state."
)
source_groups = saved["optimizer"]["param_groups"]
if len(source_groups) != len(optimizer.param_groups):
raise ValueError("Optimizer parameter groups differ.")
for source, target in zip(source_groups, optimizer.param_groups):
if (
len(source["params"]) != len(target["params"])
or source["betas"] != target["betas"]
or source["weight_decay"] != target["weight_decay"]
or source["eps"] != target["eps"]
or source.get("amsgrad", False) != target.get("amsgrad", False)
):
raise ValueError(
"Optimizer continuation requires compatible AdamW groups."
)
model.load_state_dict(saved["model"])
if optimizer is not None:
learning_rates = [group["lr"] for group in optimizer.param_groups]
optimizer.load_state_dict(saved["optimizer"])
for group, rate in zip(optimizer.param_groups, learning_rates):
group["lr"] = rate
details = {
"checkpoint": str(path.resolve()),
"source_step": saved["step"],
"source_trained_windows": saved["total_windows"],
"source_epoch": saved["epoch"],
"source_batch_cursor": saved["batch_cursor"],
"source_signatures": saved["signatures"],
"file_bytes": path.stat().st_size,
"optimizer": (
"restored AdamW moments and parameter step counters; new local scheduler and data cursors"
if optimizer is not None
else "fresh AdamW; source optimizer, scheduler and cursors are not restored"
),
"optimizer_state_count": len(saved["optimizer"]["state"])
if optimizer is not None
else 0,
"verification": "Model configuration and small dataset manifest signatures; no repeated checkpoint SHA256",
}
del saved
return details