Files
llm-atlas/experiments/k3/attnres_gradient/train.py
T
2026-07-30 07:54:38 +08:00

828 lines
30 KiB
Python

#!/usr/bin/env python3
"""Run one preregistered AttnRes activation-gradient/depth experiment cell."""
from __future__ import annotations
import argparse
import hashlib
import importlib.util
import json
import math
import os
import platform
import statistics
import sys
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Iterable
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
PROTOCOL_ID = "llm-atlas-k3-attnres-gradient-scale-v1"
ARCHITECTURES = ("baseline", "block")
DEPTHS = (16, 32)
EXPECTED_SEEDS = (2026073001, 2026073002, 2026073003)
DIAGNOSTIC_STEPS = (0, 100, 500, 2000, 4000, 8000)
FORMAL_STEPS = 8000
SMOKE_STEPS = 20
CONTEXT = 256
VOCABULARY = 256
BLOCK_GROUPS = 8
PEAK_LR = 3e-4
MIN_LR = 3e-5
WARMUP_STEPS = 400
WEIGHT_DECAY = 0.1
BETAS = (0.9, 0.95)
ADAM_EPS = 1e-8
GRAD_CLIP = 1.0
def load_round04_module() -> Any:
path = Path(__file__).resolve().parents[1] / "attnres" / "train.py"
spec = importlib.util.spec_from_file_location("k3_attnres_round04_train", path)
if spec is None or spec.loader is None:
raise RuntimeError(f"cannot import Round 04 runner from {path}")
module = importlib.util.module_from_spec(spec)
sys.modules[spec.name] = module
spec.loader.exec_module(module)
return module
round04 = load_round04_module()
def configure_round04_globals(depth: int) -> None:
round04.PROTOCOL_ID = PROTOCOL_ID
round04.LAYERS = depth
round04.SUBLAYERS = depth * 2
round04.BLOCKS = BLOCK_GROUPS
round04.SUBLAYERS_PER_BLOCK = (depth * 2) // BLOCK_GROUPS
round04.WARMUP_STEPS = WARMUP_STEPS
round04.EVAL_STEPS = DIAGNOSTIC_STEPS
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--architecture", choices=ARCHITECTURES, required=True)
parser.add_argument("--depth", type=int, choices=DEPTHS, required=True)
parser.add_argument("--seed", type=int, required=True)
parser.add_argument("--steps", type=int)
parser.add_argument("--batch-size", type=int, default=32)
parser.add_argument("--cache-dir", type=Path, required=True)
parser.add_argument("--manifest", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
parser.add_argument("--validation-windows", type=int, default=64)
parser.add_argument("--diagnostic-windows", type=int, default=16)
parser.add_argument("--eval-batch-size", type=int, default=8)
parser.add_argument("--timing-warmup", type=int, default=20)
parser.add_argument(
"--run-kind", choices=("smoke", "formal", "replay"), default="formal"
)
args = parser.parse_args()
expected_steps = SMOKE_STEPS if args.run_kind == "smoke" else FORMAL_STEPS
if args.steps is None:
args.steps = expected_steps
if args.steps != expected_steps:
raise ValueError(
f"{args.run_kind} must run exactly {expected_steps} steps, got {args.steps}"
)
if args.batch_size != 32:
raise ValueError("the frozen protocol requires batch size 32")
if args.validation_windows != 64 or args.diagnostic_windows != 16:
raise ValueError("the frozen protocol requires 64 validation / 16 diagnostic windows")
return args
def configure_determinism(seed: int) -> None:
if os.environ.get("CUBLAS_WORKSPACE_CONFIG") != ":4096:8":
raise RuntimeError("CUBLAS_WORKSPACE_CONFIG must be :4096:8 before Python starts")
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
torch.use_deterministic_algorithms(True)
torch.backends.cudnn.benchmark = False
torch.backends.cudnn.deterministic = True
torch.backends.cuda.matmul.allow_tf32 = False
torch.backends.cudnn.allow_tf32 = False
torch.set_float32_matmul_precision("highest")
def canonical_sha256(value: Any) -> str:
payload = json.dumps(
value, ensure_ascii=False, sort_keys=True, separators=(",", ":")
).encode()
return hashlib.sha256(payload).hexdigest()
def file_sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for block in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(block)
return digest.hexdigest()
def tensor_bytes(tensor: torch.Tensor) -> bytes:
value = tensor.detach().cpu().contiguous()
return (
f"{value.dtype}|{tuple(value.shape)}|".encode()
+ value.reshape(-1).view(torch.uint8).numpy().tobytes()
)
def named_state_hash(
model: nn.Module, *, include_mixers: bool | None
) -> str:
digest = hashlib.sha256()
for name, tensor in sorted(model.state_dict().items()):
is_mixer = name.startswith("mixers.") or name.startswith("output_mixer.")
if include_mixers is not None and is_mixer != include_mixers:
continue
digest.update(name.encode())
digest.update(b"\0")
digest.update(tensor_bytes(tensor))
return digest.hexdigest()
def state_structure_hash(
model: nn.Module, *, include_mixers: bool | None
) -> tuple[str, int, int]:
digest = hashlib.sha256()
tensor_count = 0
element_count = 0
for name, tensor in sorted(model.state_dict().items()):
is_mixer = name.startswith("mixers.") or name.startswith("output_mixer.")
if include_mixers is not None and is_mixer != include_mixers:
continue
digest.update(
f"{name}|{tuple(tensor.shape)}|{tensor.dtype}|{tensor.numel()}\n".encode()
)
tensor_count += 1
element_count += tensor.numel()
return digest.hexdigest(), tensor_count, element_count
def recursive_state_hash(value: Any) -> str:
digest = hashlib.sha256()
def visit(path: str, item: Any) -> None:
if torch.is_tensor(item):
digest.update(f"{path}|tensor|".encode())
digest.update(tensor_bytes(item))
elif isinstance(item, dict):
digest.update(f"{path}|dict|{len(item)}\n".encode())
for key in sorted(item, key=lambda candidate: str(candidate)):
visit(f"{path}/{key}", item[key])
elif isinstance(item, (list, tuple)):
digest.update(f"{path}|sequence|{len(item)}\n".encode())
for index, child in enumerate(item):
visit(f"{path}/{index}", child)
else:
digest.update(f"{path}|scalar|{repr(item)}\n".encode())
visit("root", value)
return digest.hexdigest()
@dataclass
class ActivationTrace:
block_outputs: list[torch.Tensor]
layer_input_rms: list[float]
branch_output_rms: list[float]
stream_state_rms: list[float]
depth_weights: list[dict[str, Any]]
output_weights: dict[str, Any] | None = None
def rms(value: torch.Tensor) -> float:
return value.float().square().mean().sqrt().detach().cpu().item()
class GradientLanguageModel(round04.ReducedLanguageModel):
"""Round 04 trunk with aligned post-MLP activation capture."""
def forward(
self, input_ids: torch.Tensor, capture: bool = False
) -> tuple[torch.Tensor, ActivationTrace | None]:
embedded = self.embed(input_ids)
trace = ActivationTrace([], [], [], [], []) if capture else None
if self.architecture == "baseline":
hidden = embedded
for block in self.blocks:
attention_input = hidden
attention_output = block.attention(block.attention_norm(attention_input))
hidden = hidden + attention_output
if trace is not None:
trace.layer_input_rms.append(rms(attention_input))
trace.branch_output_rms.append(rms(attention_output))
trace.stream_state_rms.append(rms(hidden))
mlp_input = hidden
mlp_output = block.mlp(block.mlp_norm(mlp_input))
hidden = hidden + mlp_output
if trace is not None:
hidden.retain_grad()
trace.block_outputs.append(hidden)
trace.layer_input_rms.append(rms(mlp_input))
trace.branch_output_rms.append(rms(mlp_output))
trace.stream_state_rms.append(rms(hidden))
else:
completed = [embedded]
partial: torch.Tensor | None = None
mixer_index = 0
for block in self.blocks:
for branch_index in range(2):
sources = completed + ([] if partial is None else [partial])
branch_input, weights = self.mixers[mixer_index](
sources, capture
)
mixer_index += 1
if branch_index == 0:
branch_output = block.attention(
block.attention_norm(branch_input)
)
else:
branch_output = block.mlp(block.mlp_norm(branch_input))
branch_for_residual = branch_output.float()
partial = (
branch_for_residual
if partial is None
else partial + branch_for_residual
)
if trace is not None:
trace.layer_input_rms.append(rms(branch_input))
trace.branch_output_rms.append(rms(branch_output))
trace.stream_state_rms.append(rms(partial))
trace.depth_weights.append(weights or {})
if branch_index == 1:
partial.retain_grad()
trace.block_outputs.append(partial)
if mixer_index % round04.SUBLAYERS_PER_BLOCK == 0:
completed.append(partial)
partial = None
if partial is not None or len(completed) != BLOCK_GROUPS + 1:
raise RuntimeError("Block AttnRes aggregation contract failed")
if self.output_mixer is None:
raise RuntimeError("Block AttnRes output mixer missing")
hidden, output_weights = self.output_mixer(completed, capture)
if trace is not None:
trace.output_weights = output_weights
normalized = self.final_norm(hidden)
logits = F.linear(normalized, self.token_embedding.weight)
return logits, trace
def cross_entropy(logits: torch.Tensor, targets: torch.Tensor) -> torch.Tensor:
return F.cross_entropy(
logits.float().reshape(-1, VOCABULARY), targets.reshape(-1)
)
def learning_rate(step: int, total_steps: int) -> float:
if step <= WARMUP_STEPS:
return PEAK_LR * step / WARMUP_STEPS
progress = (step - WARMUP_STEPS) / max(1, total_steps - WARMUP_STEPS)
cosine = 0.5 * (1 + math.cos(math.pi * progress))
return MIN_LR + (PEAK_LR - MIN_LR) * cosine
@torch.no_grad()
def evaluate(
model: GradientLanguageModel,
corpus: Any,
window_count: int,
eval_batch_size: int,
) -> dict[str, float]:
model.eval()
loss_sum = 0.0
target_count = 0
for begin in range(0, window_count, eval_batch_size):
end = min(begin + eval_batch_size, window_count)
inputs, targets = corpus.fixed_batch(
corpus.validation_starts, begin, end
)
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
logits, _ = model(inputs)
loss = F.cross_entropy(
logits.float().reshape(-1, VOCABULARY),
targets.reshape(-1),
reduction="sum",
)
loss_sum += loss.detach().cpu().item()
target_count += targets.numel()
nats = loss_sum / target_count
return {"cross_entropy_nats": nats, "bits_per_byte": nats / math.log(2)}
def mean(values: Iterable[float]) -> float:
return statistics.fmean(values)
def depth_statistics(values: list[float]) -> dict[str, Any]:
average = mean(values)
variance = mean((value - average) ** 2 for value in values)
quartile = len(values) // 4
first = mean(values[:quartile])
last = mean(values[-quartile:])
ratio = first / last
return {
"mean": average,
"population_cv": math.sqrt(variance) / average,
"normalized": [value / average for value in values],
"first_quartile_mean": first,
"last_quartile_mean": last,
"first_to_last_ratio": ratio,
"imbalance_abs_log_ratio": abs(math.log(ratio)),
}
def core_parameter_gradient_rms(model: GradientLanguageModel) -> list[float]:
values = []
for block in model.blocks:
sum_square = 0.0
count = 0
for parameter in block.parameters():
if parameter.grad is None:
raise RuntimeError("missing core parameter gradient")
gradient = parameter.grad.detach().float()
if not torch.isfinite(gradient).all():
raise RuntimeError("non-finite core parameter gradient")
sum_square += gradient.square().sum().detach().cpu().item()
count += gradient.numel()
values.append(math.sqrt(sum_square / count))
return values
def activation_storage_unique(outputs: list[torch.Tensor]) -> bool:
pointers = [output.untyped_storage().data_ptr() for output in outputs]
return len(pointers) == len(set(pointers))
def diagnostic(
model: GradientLanguageModel,
corpus: Any,
window_count: int,
*,
loss_scale: float = 1.0,
) -> dict[str, Any]:
model.eval()
model.zero_grad(set_to_none=True)
inputs, targets = corpus.fixed_batch(
corpus.diagnostic_starts, 0, window_count
)
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
logits, trace = model(inputs, capture=True)
unscaled_loss = cross_entropy(logits, targets)
loss = unscaled_loss * loss_scale
if trace is None or len(trace.block_outputs) != len(model.blocks):
raise RuntimeError("aligned activation capture count mismatch")
expected_shape = (window_count, CONTEXT, round04.D_MODEL)
if any(tuple(output.shape) != expected_shape for output in trace.block_outputs):
raise RuntimeError("aligned activation capture shape mismatch")
if any(output.dtype != torch.float32 for output in trace.block_outputs):
raise RuntimeError("aligned activation capture must use FP32 residual state")
if not activation_storage_unique(trace.block_outputs):
raise RuntimeError("captured block outputs alias storage")
loss.backward()
activation_grad_rms = []
activation_output_rms = []
activation_dtypes = []
for output in trace.block_outputs:
if output.grad is None:
raise RuntimeError("captured activation gradient is None")
gradient = output.grad.detach().float()
if not torch.isfinite(gradient).all():
raise RuntimeError("captured activation gradient is non-finite")
activation_grad_rms.append(
gradient.square().mean().sqrt().detach().cpu().item()
)
activation_output_rms.append(rms(output))
activation_dtypes.append(str(output.dtype))
parameter_grad_rms = core_parameter_gradient_rms(model)
return {
"loss_nats": unscaled_loss.detach().cpu().item(),
"bits_per_byte": unscaled_loss.detach().cpu().item() / math.log(2),
"loss_scale": loss_scale,
"capture": {
"count": len(trace.block_outputs),
"shape": list(expected_shape),
"dtypes": activation_dtypes,
"all_gradients_finite": True,
"all_gradients_present": True,
"storage_unique": True,
"position": (
"post-MLP Transformer-block output; Block AttnRes is captured "
"before aggregation-partial reset"
),
},
"activation_grad_rms_by_block": activation_grad_rms,
"activation_grad_statistics": depth_statistics(activation_grad_rms),
"activation_output_rms_by_block": activation_output_rms,
"activation_output_statistics": depth_statistics(activation_output_rms),
"core_parameter_grad_rms_by_block": parameter_grad_rms,
"core_parameter_grad_statistics": depth_statistics(parameter_grad_rms),
"layer_input_rms_by_sublayer": trace.layer_input_rms,
"branch_output_rms_by_sublayer": trace.branch_output_rms,
"stream_state_rms_by_sublayer": trace.stream_state_rms,
"depth_weights": trace.depth_weights,
"output_weights": trace.output_weights,
}
def loss_scale_gate(
model: GradientLanguageModel, corpus: Any, window_count: int
) -> tuple[dict[str, Any], dict[str, Any]]:
base = diagnostic(model, corpus, window_count, loss_scale=1.0)
doubled = diagnostic(model, corpus, window_count, loss_scale=2.0)
base_values = base["activation_grad_rms_by_block"]
doubled_values = doubled["activation_grad_rms_by_block"]
ratios = [
doubled_value / base_value
for base_value, doubled_value in zip(base_values, doubled_values)
]
base_stats = base["activation_grad_statistics"]
doubled_stats = doubled["activation_grad_statistics"]
cv_delta = abs(
doubled_stats["population_cv"] - base_stats["population_cv"]
)
ratio_delta = abs(
doubled_stats["first_to_last_ratio"]
- base_stats["first_to_last_ratio"]
)
normalized_max_delta = max(
abs(left - right)
for left, right in zip(
base_stats["normalized"], doubled_stats["normalized"]
)
)
passed = (
all(abs(ratio - 2.0) <= 1e-5 for ratio in ratios)
and cv_delta <= 1e-6
and ratio_delta <= 1e-6
and normalized_max_delta <= 1e-6
)
gate = {
"passed": passed,
"per_block_scale_ratios": ratios,
"max_abs_scale_ratio_error": max(abs(ratio - 2.0) for ratio in ratios),
"population_cv_abs_delta": cv_delta,
"first_to_last_ratio_abs_delta": ratio_delta,
"normalized_spectrum_max_abs_delta": normalized_max_delta,
"thresholds": {
"scale_ratio_abs": 1e-5,
"shape_abs": 1e-6,
},
}
if not passed:
raise RuntimeError(f"loss-scale diagnostic gate failed: {gate}")
return base, gate
def percentile(values: list[float], quantile: float) -> float:
return float(np.quantile(np.asarray(values, dtype=np.float64), quantile))
def parameter_inventory(model: GradientLanguageModel) -> dict[str, int]:
total = sum(parameter.numel() for parameter in model.parameters())
mixer = sum(
parameter.numel()
for name, parameter in model.named_parameters()
if name.startswith("mixers.") or name.startswith("output_mixer.")
)
return {
"total": total,
"core": total - mixer,
"mixer": mixer,
"embedding": (
model.token_embedding.weight.numel()
+ model.position_embedding.weight.numel()
),
}
def model_input_gate_hashes(
corpus: Any, manifest: dict[str, Any], seed: int, batch_size: int
) -> dict[str, str]:
values: dict[str, str] = {}
for step in manifest["windows"]["gate_steps"]:
raw_digest = hashlib.sha256()
for row in range(batch_size):
start = round04.window_start(seed, step, row, len(corpus.train))
raw_digest.update(
np.asarray(
corpus.train[start : start + CONTEXT + 1], dtype=np.uint8
).tobytes()
)
expected_raw_hash = manifest["windows"][
"gate_training_tensor_sha256"
][str(seed)][str(step)]
if raw_digest.hexdigest() != expected_raw_hash:
raise RuntimeError(f"manifest gate tensor mismatch at step {step}")
inputs, targets = corpus.training_batch(seed, step, batch_size)
digest = hashlib.sha256()
digest.update(tensor_bytes(inputs))
digest.update(tensor_bytes(targets))
values[str(step)] = digest.hexdigest()
return values
def main() -> None:
args = parse_args()
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required by the frozen protocol")
if args.seed not in EXPECTED_SEEDS:
raise ValueError(f"seed is not preregistered: {args.seed}")
configure_round04_globals(args.depth)
configure_determinism(args.seed)
device = torch.device("cuda")
manifest = json.loads(args.manifest.read_text())
if manifest["protocol_id"] != PROTOCOL_ID:
raise ValueError("manifest protocol mismatch")
if manifest["windows"]["formal_steps"] != FORMAL_STEPS:
raise ValueError("manifest formal-step mismatch")
corpus = round04.ByteCorpus(args.cache_dir, manifest, device)
model = GradientLanguageModel(args.architecture).to(device)
public_structure_hash, public_tensors, public_elements = state_structure_hash(
model, include_mixers=False
)
initial_public_hash = named_state_hash(model, include_mixers=False)
initial_mixer_hash = (
named_state_hash(model, include_mixers=True)
if args.architecture == "block"
else None
)
input_gate_hashes = model_input_gate_hashes(
corpus, manifest, args.seed, args.batch_size
)
decay_parameters: list[nn.Parameter] = []
no_decay_parameters: list[nn.Parameter] = []
for parameter in model.parameters():
(decay_parameters if parameter.ndim >= 2 else no_decay_parameters).append(
parameter
)
optimizer = torch.optim.AdamW(
[
{"params": decay_parameters, "weight_decay": WEIGHT_DECAY},
{"params": no_decay_parameters, "weight_decay": 0.0},
],
lr=PEAK_LR,
betas=BETAS,
eps=ADAM_EPS,
)
evaluation_steps = sorted(
set(step for step in DIAGNOSTIC_STEPS if step <= args.steps)
| {0, args.steps}
)
evaluations = [
{
"step": 0,
**evaluate(
model, corpus, args.validation_windows, args.eval_batch_size
),
}
]
if args.run_kind == "smoke":
initial_diagnostic, gradient_gate = loss_scale_gate(
model, corpus, args.diagnostic_windows
)
else:
initial_diagnostic = diagnostic(
model, corpus, args.diagnostic_windows
)
gradient_gate = None
diagnostics = [{"step": 0, **initial_diagnostic}]
print(
json.dumps(
{
"event": "diagnostic",
"step": 0,
"architecture": args.architecture,
"depth": args.depth,
"validation_bpc": evaluations[0]["bits_per_byte"],
"activation_gradient_cv": initial_diagnostic[
"activation_grad_statistics"
]["population_cv"],
},
sort_keys=True,
),
flush=True,
)
model.zero_grad(set_to_none=True)
training_history: list[dict[str, float | int]] = []
step_times: list[float] = []
model.train()
for step in range(1, args.steps + 1):
lr = learning_rate(step, args.steps)
for group in optimizer.param_groups:
group["lr"] = lr
inputs, targets = corpus.training_batch(args.seed, step, args.batch_size)
optimizer.zero_grad(set_to_none=True)
torch.cuda.synchronize()
started = time.perf_counter()
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
logits, _ = model(inputs)
loss = cross_entropy(logits, targets)
if not torch.isfinite(loss):
raise RuntimeError(f"non-finite loss at step {step}: {loss}")
loss.backward()
unclipped_norm = torch.nn.utils.clip_grad_norm_(
model.parameters(), GRAD_CLIP
)
optimizer.step()
torch.cuda.synchronize()
elapsed_ms = (time.perf_counter() - started) * 1000
if step == args.timing_warmup:
torch.cuda.reset_peak_memory_stats()
elif step > args.timing_warmup:
step_times.append(elapsed_ms)
if step == 1 or step % 10 == 0 or step == args.steps:
training_history.append(
{
"step": step,
"loss_nats": loss.detach().cpu().item(),
"bits_per_byte": loss.detach().cpu().item() / math.log(2),
"learning_rate": lr,
"unclipped_grad_norm": float(unclipped_norm.detach().cpu()),
}
)
if step in evaluation_steps and step != 0:
evaluations.append(
{
"step": step,
**evaluate(
model,
corpus,
args.validation_windows,
args.eval_batch_size,
),
}
)
diagnostics.append(
{
"step": step,
**diagnostic(
model, corpus, args.diagnostic_windows
),
}
)
print(
json.dumps(
{
"event": "diagnostic",
"step": step,
"architecture": args.architecture,
"depth": args.depth,
"validation_bpc": evaluations[-1]["bits_per_byte"],
"activation_gradient_cv": diagnostics[-1][
"activation_grad_statistics"
]["population_cv"],
},
sort_keys=True,
),
flush=True,
)
model.zero_grad(set_to_none=True)
model.train()
training_peak_allocated = torch.cuda.max_memory_allocated()
training_peak_reserved = torch.cuda.max_memory_reserved()
final_public_hash = named_state_hash(model, include_mixers=False)
final_mixer_hash = (
named_state_hash(model, include_mixers=True)
if args.architecture == "block"
else None
)
final_full_hash = named_state_hash(model, include_mixers=None)
optimizer_hash = recursive_state_hash(optimizer.state_dict())
timing = {
"warmup_steps_excluded": args.timing_warmup,
"measured_steps": len(step_times),
"mean_ms": mean(step_times) if step_times else None,
"median_ms": statistics.median(step_times) if step_times else None,
"p95_ms": percentile(step_times, 0.95) if step_times else None,
"peak_allocated_bytes": training_peak_allocated,
"peak_reserved_bytes": training_peak_reserved,
}
result = {
"schema_version": 1,
"protocol_id": PROTOCOL_ID,
"run_kind": args.run_kind,
"architecture": args.architecture,
"depth": args.depth,
"seed": args.seed,
"steps": args.steps,
"batch_size": args.batch_size,
"target_bytes_seen": args.steps * args.batch_size * CONTEXT,
"manifest": {
"path": str(args.manifest),
"file_sha256": file_sha256(args.manifest),
"formal_schedule_sha256": manifest["windows"][
"formal_schedule_sha256"
],
"validation_tensor_sha256": manifest["windows"][
"validation_tensor_sha256"
],
"diagnostic_tensor_sha256": manifest["windows"][
"diagnostic_tensor_sha256"
],
"input_gate_tensor_hashes": input_gate_hashes,
},
"model": {
"layers": args.depth,
"sublayers": args.depth * 2,
"attnres_aggregation_groups": BLOCK_GROUPS,
"sublayers_per_attnres_group": args.depth * 2 // BLOCK_GROUPS,
"transformer_blocks_per_attnres_group": args.depth // BLOCK_GROUPS,
"d_model": round04.D_MODEL,
"heads": round04.HEADS,
"d_head": round04.D_HEAD,
"d_ff": round04.D_FF,
"context": CONTEXT,
"vocabulary": VOCABULARY,
"parameters": parameter_inventory(model),
},
"optimizer": {
"name": "AdamW",
"betas": list(BETAS),
"epsilon": ADAM_EPS,
"weight_decay_ndim_ge_2": WEIGHT_DECAY,
"peak_lr": PEAK_LR,
"min_lr": MIN_LR,
"warmup_steps": WARMUP_STEPS,
"grad_clip": GRAD_CLIP,
},
"hashes": {
"initial_public_parameter_structure": public_structure_hash,
"initial_public_parameter_tensors": public_tensors,
"initial_public_parameter_elements": public_elements,
"initial_public_parameters": initial_public_hash,
"initial_mixer_parameters": initial_mixer_hash,
"final_public_parameters": final_public_hash,
"final_mixer_parameters": final_mixer_hash,
"final_model_state": final_full_hash,
"final_optimizer_state": optimizer_hash,
},
"evaluations": evaluations,
"diagnostics": diagnostics,
"training_history": training_history,
"gradient_gate": gradient_gate,
"timing": timing,
"environment": {
"python": platform.python_version(),
"torch": torch.__version__,
"cuda": torch.version.cuda,
"gpu": torch.cuda.get_device_name(0),
"compute_capability": list(torch.cuda.get_device_capability(0)),
"cublas_workspace_config": os.environ["CUBLAS_WORKSPACE_CONFIG"],
"deterministic_algorithms": torch.are_deterministic_algorithms_enabled(),
"autocast": "cuda-bfloat16-forward-fp32-cross-entropy",
"compile": False,
},
}
result["canonical_sha256_without_self"] = canonical_sha256(result)
args.output.parent.mkdir(parents=True, exist_ok=True)
temporary = args.output.with_suffix(args.output.suffix + ".tmp")
temporary.write_text(
json.dumps(result, ensure_ascii=False, indent=2, sort_keys=True) + "\n"
)
os.replace(temporary, args.output)
print(
json.dumps(
{
"output": str(args.output),
"run_kind": args.run_kind,
"architecture": args.architecture,
"depth": args.depth,
"seed": args.seed,
"steps": args.steps,
"final_bpc": evaluations[-1]["bits_per_byte"],
"final_activation_gradient_cv": diagnostics[-1][
"activation_grad_statistics"
]["population_cv"],
"canonical_sha256": result["canonical_sha256_without_self"],
"timing": timing,
},
ensure_ascii=False,
indent=2,
)
)
if __name__ == "__main__":
main()