stackcraft-clef-flash-lora / code /scripts /probe_clef_training.py
nima1's picture
Publish verified checkpoint and losslessly compressed study evidence
4be6a52 verified
Raw History Blame Contribute Delete
10.5 kB
"""Bounded real-weight GPU training/reload gate; never manages other workloads.
Run from the repository with uv run --locked --extra ml python scripts/probe_clef_training.py.
The parent process must impose a wall-clock timeout and restore borrowed GPU services.
"""
import argparse
import gc
import hashlib
import json
import time
from pathlib import Path
import torch
from stackcraft.clef import ClefPlayer, encode_observation
from stackcraft.data import audit_dataset
from stackcraft.players import observe
from stackcraft.provenance import source_identity
from stackcraft.schema import GameState
from stackcraft.training import (
decision_loss,
load_checkpoint,
parameter_hashes,
prepare_trainable,
save_checkpoint,
)
def observation(row):
raw = row["observation"]
return observe(
GameState(tuple(tuple(r) for r in raw["board"]), 0, 0, raw["current"], raw["next_piece"])
)
def write_json(path, value):
path.write_text(json.dumps(value, indent=2, allow_nan=False) + "\n")
def probabilities(player, rows):
return [player.choose(observation(row)).probabilities for row in rows]
def train_steps(player, rows, *, steps, learning_rate):
model = player.model
model.train()
if model._stackcraft_training["mode"] == "head":
model.language_model.eval()
parameters = [p for p in model.parameters() if p.requires_grad]
optimizer = torch.optim.AdamW(parameters, lr=learning_rate)
events = []
for index in range(steps):
row = rows[index % len(rows)]
encoded = encode_observation(
observation(row), player.processor.tokenizer, player.native, player.max_length
)
batch = player.native.collate_records(
[encoded], player.processor.tokenizer.pad_token_id, torch.device("cuda")
)
optimizer.zero_grad(set_to_none=True)
started = time.monotonic()
logits = model(batch)[0][0]
loss = decision_loss(logits, encoded, row["action_id"])
if not torch.isfinite(loss):
raise RuntimeError("training loss is nonfinite")
loss.backward()
gradient_sums = {"head": 0.0, "lora": 0.0}
for name, param in model.named_parameters():
if param.grad is None:
continue
if not torch.isfinite(param.grad).all():
raise RuntimeError(f"nonfinite gradient: {name}")
group = "lora" if "lora_" in name else "head"
gradient_sums[group] += float(param.grad.detach().abs().sum())
if gradient_sums["head"] <= 0:
raise RuntimeError("no nonzero decision-head gradients")
if any("lora_" in name for name, p in model.named_parameters() if p.requires_grad):
if gradient_sums["lora"] <= 0:
raise RuntimeError("no nonzero LoRA gradients")
torch.nn.utils.clip_grad_norm_(parameters, 1.0, error_if_nonfinite=True)
optimizer.step()
torch.cuda.synchronize()
event = {
"step": index + 1,
"row_id": row["id"],
"tokens": len(encoded.input_ids),
"loss": float(loss.detach()),
"gradient_abs_sums": gradient_sums,
"seconds": time.monotonic() - started,
"peak_allocated_bytes": torch.cuda.max_memory_allocated(),
"peak_reserved_bytes": torch.cuda.max_memory_reserved(),
}
events.append(event)
print(json.dumps(event), flush=True)
del optimizer
model.zero_grad(set_to_none=True)
model.eval()
torch.cuda.empty_cache()
return events
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--output", type=Path, required=True)
parser.add_argument("--dataset", type=Path, default=Path("data/study-v1"))
parser.add_argument("--reload", type=Path)
args = parser.parse_args()
if args.output.exists():
parser.error("output already exists; choose a new directory")
args.output.mkdir(parents=True)
started = time.monotonic()
report = {"status": "running", "reload": bool(args.reload)}
write_json(args.output / "report.json", report)
try:
if not torch.cuda.is_available():
raise RuntimeError("CUDA unavailable; this gate requires real GPU training")
free, total = torch.cuda.mem_get_info()
if free < 25 * 1024**3:
raise RuntimeError(f"requires at least25GiB free before loading; available={free}")
torch.manual_seed(42)
torch.set_num_threads(8)
torch.backends.cuda.matmul.allow_tf32 = False
manifest_path = args.dataset / "manifest.json"
manifest = json.loads(manifest_path.read_text())
records = {
split: [
json.loads(line)
for line in (args.dataset / f"{split}.jsonl").read_text().splitlines()
]
for split in ("train", "validation")
}
audit_dataset(records, manifest)
# Development probe uses only the first four training positions, never test seeds.
rows = records["train"][:4]
report.update(
gpu=torch.cuda.get_device_name(),
total_vram=total,
initial_free_vram=free,
torch=torch.__version__,
dataset_manifest_sha256=hashlib.sha256(manifest_path.read_bytes()).hexdigest(),
row_ids=[row["id"] for row in rows],
)
report.update(source_identity(Path(__file__).resolve().parents[1]))
player = ClefPlayer.from_pretrained(trust_pinned_code=True)
torch.cuda.reset_peak_memory_stats()
if args.reload:
reference = json.loads((args.reload / "reference.json").read_text())
for key in ("row_ids", "dataset_manifest_sha256"):
if reference.get(key) != report[key]:
raise RuntimeError(f"reload reference {key} differs")
player.model = load_checkpoint(player.model, args.reload)
actual = probabilities(player, rows)
if len(reference["probabilities"]) != len(actual):
raise RuntimeError("reference record count differs")
delta = 0.0
for expected, observed in zip(reference["probabilities"], actual, strict=True):
if expected.keys() != observed.keys():
raise RuntimeError("reload probability option set differs")
delta = max(delta, *(abs(expected[k] - observed[k]) for k in expected))
if delta > 1e-4:
raise RuntimeError(f"fresh-process probability drift{delta} exceeds1e-4")
report.update(max_absolute_probability_difference=delta, tolerance=1e-4)
else:
native_probabilities = probabilities(player, rows)
report["native_probabilities"] = native_probabilities
report["runtime_config"] = player.runtime_config
report["load_and_native_seconds"] = time.monotonic() - started
write_json(args.output / "report.json", report)
prepare_trainable(player.model, mode="head")
wrapped = probabilities(player, rows)
report["fp32_head_initial_max_probability_drift"] = max(
abs(a[key] - b[key])
for a, b in zip(native_probabilities, wrapped, strict=True)
for key in a
)
before_head = parameter_hashes(player.model, trainable=True)
before_frozen = parameter_hashes(player.model, trainable=False)
report["head_steps"] = train_steps(player, rows, steps=3, learning_rate=1e-5)
if before_head == parameter_hashes(player.model, trainable=True):
raise RuntimeError("head parameters did not change")
if before_frozen != parameter_hashes(player.model, trainable=False):
raise RuntimeError("frozen parameters changed during head training")
save_checkpoint(
player.model,
args.output / "head-checkpoint",
extra_metadata={"probe_rows": report["row_ids"]},
)
write_json(
args.output / "head-checkpoint" / "reference.json",
{
"probabilities": probabilities(player, rows),
"row_ids": report["row_ids"],
"dataset_manifest_sha256": report["dataset_manifest_sha256"],
},
)
report["head_frozen_parameters_unchanged"] = True
write_json(args.output / "report.json", report)
del player
gc.collect()
torch.cuda.empty_cache()
player = ClefPlayer.from_pretrained(trust_pinned_code=True)
prepare_trainable(player.model, mode="lora", rank=4)
before_lora = parameter_hashes(player.model, trainable=True)
before_frozen = parameter_hashes(player.model, trainable=False)
report["lora_steps"] = train_steps(player, rows, steps=5, learning_rate=1e-5)
after_lora = parameter_hashes(player.model, trainable=True)
if not any(
"lora_" in name and value != after_lora[name] for name, value in before_lora.items()
):
raise RuntimeError("LoRA parameters did not change")
if before_frozen != parameter_hashes(player.model, trainable=False):
raise RuntimeError("frozen parameters changed during LoRA training")
checkpoint = args.output / "checkpoint"
save_checkpoint(
player.model, checkpoint, extra_metadata={"probe_rows": report["row_ids"]}
)
write_json(
checkpoint / "reference.json",
{
"probabilities": probabilities(player, rows),
"row_ids": report["row_ids"],
"dataset_manifest_sha256": report["dataset_manifest_sha256"],
},
)
report["checkpoint"] = str(checkpoint)
report["frozen_parameters_unchanged"] = True
report.update(status="passed", elapsed_seconds=time.monotonic() - started)
except Exception as error:
report.update(status="failed", error=f"{type(error).__name__}: {error}")
raise
finally:
report["elapsed_seconds"] = time.monotonic() - started
write_json(args.output / "report.json", report)
if __name__ == "__main__":
main()