vimeml-tiny-ja-v2.1 / source /scripts /training /check_model_v2.py
Voltline's picture
Release VimeML V2.1 step40000 FP32 and Core ML INT8 (GPL-2.0)
29f25be verified
Raw History Blame Contribute Delete
5.46 kB
"""One compact V2 architecture acceptance run; no corpus training or W&B run."""
import argparse
import json
import sys
import time
import tomllib
from pathlib import Path
ROOT = Path(__file__).resolve().parents[2]
sys.path.insert(0, str(ROOT / "src"))
import torch
from torch.nn import functional as F
from vimeml.training.model_factory import (
checkpoint_format,
configuration_for,
create_model,
model_from_checkpoint,
)
from vimeml.training.train import atomic_checkpoint, optimizer_for, validate_config
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--config", type=Path, default=ROOT / "configs/train-v2.toml")
parser.add_argument(
"--output", type=Path, default=ROOT / "outputs/model-checks/tiny-ja-v2-phase-b"
)
args = parser.parse_args()
config = tomllib.loads(args.config.read_text(encoding="utf-8"))
architecture = config["architecture"]
model_config, settings = validate_config(config)
device = torch.device(settings["device"])
torch.set_num_threads(settings["cpu_threads"])
torch.manual_seed(settings["seed"])
started = time.perf_counter()
model = create_model(architecture, model_config).to(device)
assert model.parameter_count() == 12_537_920
assert model.token_embedding.weight is model.lm_head.weight
assert all(
layer.bias is None for layer in model.modules() if isinstance(layer, torch.nn.Linear)
)
inputs = torch.randint(
4, model_config.vocab_size, (4, model_config.context_length), device=device
)
labels = inputs.roll(-1, dims=1)
labels[:, 96:] = -100
model.eval()
with torch.no_grad():
logits = model(inputs)
assert logits.shape == (4, 128, 16384)
expected_loss = F.cross_entropy(
logits.flatten(0, 1), labels.flatten(), ignore_index=-100, reduction="sum"
)
result = model(inputs, labels)
torch.testing.assert_close(result["loss_sum"], expected_loss, rtol=1e-5, atol=1e-4)
changed = inputs.clone()
changed[:, 64:] = (changed[:, 64:] + 7) % model_config.vocab_size
causal_logits = model(changed)
causal_error = float((logits[:, :64] - causal_logits[:, :64]).abs().max())
assert causal_error == 0.0
del logits, causal_logits
model.train()
optimizer = optimizer_for(model, settings, device)
optimizer.zero_grad(set_to_none=True)
with torch.autocast(device_type=device.type, dtype=torch.bfloat16):
result = model(inputs, labels)
loss = result["loss_sum"] / result["token_count"]
loss.backward()
grad_norm = torch.nn.utils.clip_grad_norm_(
model.parameters(), settings["grad_clip"], error_if_nonfinite=True
)
assert torch.isfinite(loss)
assert all(torch.isfinite(p.grad).all() for p in model.parameters() if p.grad is not None)
optimizer.step()
optimizer.zero_grad(set_to_none=True)
args.output.mkdir(parents=True, exist_ok=True)
checkpoint_path = args.output / "one-update.pt"
atomic_checkpoint(
checkpoint_path,
{
"format": checkpoint_format(architecture),
"architecture": architecture,
"model": model.state_dict(),
"model_config": model.configuration(),
"optimizer": optimizer.state_dict(),
"config": config,
"step": 1,
"purpose": "Architecture acceptance only; random initialization plus one synthetic update.",
},
)
saved = torch.load(checkpoint_path, map_location=device, weights_only=True)
restored = model_from_checkpoint(saved).to(device).eval()
model.eval()
assert restored.token_embedding.weight is restored.lm_head.weight
with torch.no_grad():
reference = model(inputs[:1])
actual = restored(inputs[:1])
torch.testing.assert_close(actual, reference, rtol=0, atol=0)
restored_optimizer = optimizer_for(restored, settings, device)
restored_optimizer.load_state_dict(saved["optimizer"])
assert len(restored_optimizer.state) == len(optimizer.state)
# The default architecture continues to construct the frozen V1 model.
v1 = create_model("tiny_gpt_v1", configuration_for("tiny_gpt_v1", {}))
assert v1.parameter_count() == 7_386_624
report = {
"status": "passed",
"architecture": architecture,
"parameters": model.parameter_count(),
"model": model.configuration(),
"device": str(device),
"precision": "bf16",
"torch_version": str(torch.__version__),
"forward_shape": [4, 128, 16384],
"causal_prefix_max_error": causal_error,
"masked_loss_matches_full_logits": True,
"weight_tying": True,
"linear_bias_count": 0,
"loss_one_update": float(loss.detach()),
"grad_norm": float(grad_norm),
"checkpoint_roundtrip_exact": True,
"optimizer_state_restored": True,
"v1_parameters": v1.parameter_count(),
"elapsed_seconds": time.perf_counter() - started,
"validation_scope": "One synthetic optimizer update; not corpus training or an accuracy evaluation.",
"deployment_validated": False,
}
(args.output / "report.json").write_text(
json.dumps(report, ensure_ascii=False, indent=2) + "\n", encoding="utf-8"
)
print(json.dumps(report, ensure_ascii=False, indent=2))
if __name__ == "__main__":
main()