TinyMOSS-Diarize / runtime /quantlib /checkpoint_audit.py
yongyizang's picture
TinyMOSS-Diarize: 2.911-bit packed weights, runtime, and model card
7ccb33d verified
Raw
History Blame Contribute Delete
4.8 kB
"""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