File size: 5,464 Bytes
29f25be | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 | """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()
|