"""Strict topology/source audit for KD and VESPO student checkpoints.""" from __future__ import annotations import json from pathlib import Path from typing import Iterable import torch SUPPORTED_CHECKPOINT_FORMATS = { "moss-sherry-kd-qat-v1", "moss-sherry-vespo-v1", # This bounded pilot keeps optimizer/head state in the same document, but # its ``student`` view and topology config are standalone and deployable. "moss-speaker-state-pilot-checkpoint-v1", } def audit_quantized_student_state( state: dict[str, torch.Tensor], *, e2e: bool, required_quantized_keys: Iterable[str] = (), ) -> None: """Validate required FP32 masters and the tied embedding checkpoint view.""" if not all(isinstance(key, str) and torch.is_tensor(value) for key, value in state.items()): raise ValueError("checkpoint student state must map string keys to tensors") required = set(required_quantized_keys) missing = sorted(required - set(state)) if missing: raise ValueError(f"checkpoint is missing quantized tensors: {missing[:20]}") rounded = sorted(key for key in required if state[key].dtype != torch.float32) if rounded: details = [(key, str(state[key].dtype)) for key in rounded[:20]] raise ValueError(f"quantized master tensors must be FP32: {details}") embed_key = "model.language_model.embed_tokens.weight" head_key = "lm_head.weight" if e2e: if embed_key not in state or head_key not in state: raise ValueError("checkpoint is missing tied embedding/lm_head weights") if not torch.equal(state[embed_key], state[head_key]): raise ValueError("checkpoint embedding and lm_head weights disagree") def audited_student_checkpoint( path: Path, *, e2e: bool, embed_bits: int, embed_group_size: int, required_quantized_keys: Iterable[str] = (), ) -> tuple[dict[str, torch.Tensor], str, int, str | None]: document = torch.load(path, map_location="cpu", weights_only=False) bare_snapshot = False if isinstance(document, dict) and isinstance(document.get("student"), dict): state = document["student"] topology = document.get("config") version = str(document.get("version") or f"kd-{int(document.get('step', 0))}") step = int(document.get("step", 0)) checkpoint_format = document.get("format") if checkpoint_format not in SUPPORTED_CHECKPOINT_FORMATS: raise ValueError(f"unsupported student checkpoint format: {checkpoint_format!r}") elif isinstance(document, dict) and all(torch.is_tensor(value) for value in document.values()): bare_snapshot = True state = document metadata_path = path.with_name("metadata.json") if not metadata_path.is_file(): raise ValueError("bare VESPO snapshot requires sibling metadata.json") topology = json.loads(metadata_path.read_text(encoding="utf-8")) if topology.get("schema") != "vespo-snapshot-v1": raise ValueError("bare VESPO snapshot metadata schema is invalid") if topology.get("dtype") != "mixed_fp32_quant_bf16_residual": raise ValueError( f"legacy/rounded VESPO snapshot dtype is not allowed: {topology.get('dtype')!r}" ) version = str(topology.get("version") or path.parent.name) step = int(topology.get("step", 0)) checkpoint_format = "vespo-snapshot-v1" else: raise ValueError(f"unsupported student checkpoint format: {path}") if not isinstance(topology, dict): raise ValueError("checkpoint is missing topology config") recorded_e2e = topology.get("e2e") if type(recorded_e2e) is not bool or recorded_e2e != e2e: raise ValueError(f"checkpoint e2e topology mismatch: {topology.get('e2e')!r} != {e2e}") recorded_bits = topology.get("embed_bits") recorded_group = topology.get("embed_group_size") if e2e: # The original W4 e2e trainer predates explicit embed metadata. Missing # values are unambiguous only for that legacy W4 full checkpoint. if not (embed_bits == 4 and not bare_snapshot and recorded_bits is None): if recorded_bits != embed_bits: raise ValueError( f"checkpoint embedding precision mismatch: {recorded_bits!r} != {embed_bits}" ) if embed_bits == 3 and recorded_group != embed_group_size: raise ValueError( f"checkpoint embedding group mismatch: {recorded_group!r} != {embed_group_size}" ) audit_quantized_student_state( state, e2e=e2e, required_quantized_keys=required_quantized_keys, ) return state, version, step, checkpoint_format