from __future__ import annotations import json import os from pathlib import Path from configs import cfg from src.kd_contracts import ( PROVENANCE_SCHEMA_VERSION, build_shard_schema, canonical_revision, collect_model_vocab_sizes, sha256_file, ) def resolve_model_vocab_size(model, tokenizer, label: str, log) -> int: model_sizes = collect_model_vocab_sizes(model) if not model_sizes: log.error(f"{label} model does not expose a usable vocab size view.") raise SystemExit(1) unique_sizes = sorted(set(model_sizes.values())) if len(unique_sizes) != 1: details = ", ".join(f"{name}={size:,}" for name, size in sorted(model_sizes.items())) log.error(f"{label} vocab mismatch across checkpoint artifacts: {details}") raise SystemExit(1) model_vocab_size = unique_sizes[0] tokenizer_vocab_size = len(tokenizer) if model_vocab_size < tokenizer_vocab_size: log.error( f"{label} tokenizer length ({tokenizer_vocab_size:,}) exceeds " f"the model vocab size ({model_vocab_size:,})." ) log.error("Repair or regenerate the checkpoint before using it for distillation.") raise SystemExit(1) if model_vocab_size > tokenizer_vocab_size: log.info( f" {label} model vocab is padded beyond the tokenizer range: " f"tokenizer={tokenizer_vocab_size:,}, model={model_vocab_size:,}" ) return model_vocab_size def validate_provenance( prov_path: str, data_path: str, dataset, teacher_tokenizer_contract: dict, student_tokenizer_contract: dict, log, ) -> None: if not os.path.exists(prov_path): log.error("Missing _provenance.json in the logits directory.") log.error("Regenerate the current teacher-logit shard metadata.") raise SystemExit(1) with open(prov_path, "r", encoding="utf-8") as f: prov = json.load(f) schema_version = prov.get("schema_version") if schema_version != PROVENANCE_SCHEMA_VERSION: log.error( f"Unsupported shard provenance schema: {schema_version!r}. " f"Expected {PROVENANCE_SCHEMA_VERSION}." ) log.error("Regenerate the teacher-logit shards.") raise SystemExit(1) teacher_meta = prov.get("teacher", {}) student_meta = prov.get("student", {}) current_data_sha = sha256_file(data_path) actual_shard_count = sum(1 for _ in Path(prov_path).parent.glob("shard_*.pt")) provenance_num_samples = prov.get("num_samples") try: provenance_num_samples_int = int(provenance_num_samples) except (TypeError, ValueError): log.error(f"PROVENANCE MISMATCH: num_samples is {provenance_num_samples!r}, expected an integer.") log.error("Regenerate compatible teacher-logit shards.") raise SystemExit(1) if provenance_num_samples_int < len(dataset): log.error( f"PROVENANCE MISMATCH: num_samples is {provenance_num_samples_int}, " f"but the requested dataset has {len(dataset)} samples." ) log.error("Regenerate compatible teacher-logit shards.") raise SystemExit(1) if provenance_num_samples_int > len(dataset): log.warning( f" Provenance contains {provenance_num_samples_int:,} samples; " f"training is using the first {len(dataset):,}. This is expected for smoke tests." ) expected_pairs = [ ("shard_count", prov.get("shard_count"), actual_shard_count), ("samples_per_shard", prov.get("samples_per_shard"), dataset.samples_per_shard), ("data_sha256", prov.get("data_sha256"), current_data_sha), ("max_seq_len", prov.get("max_seq_len"), cfg.data.max_seq_len), ("top_k", prov.get("top_k"), cfg.training.top_k), ("temperature", prov.get("temperature"), float(cfg.training.temperature)), ("teacher.model", teacher_meta.get("model"), cfg.model.teacher), ( "teacher.revision", teacher_meta.get("revision"), canonical_revision(cfg.model.teacher_revision), ), ( "teacher.tokenizer_size", teacher_meta.get("tokenizer_size"), teacher_tokenizer_contract["full_vocab_size"], ), ( "teacher.tokenizer_fingerprint", teacher_meta.get("tokenizer_fingerprint"), teacher_tokenizer_contract["fingerprint"], ), ("student.model", student_meta.get("model"), getattr(cfg.model, "tokenizer", cfg.model.student)), ( "student.revision", student_meta.get("revision"), canonical_revision(getattr(cfg.model, "tokenizer_revision", cfg.model.student_revision)), ), ( "student.tokenizer_size", student_meta.get("tokenizer_size"), student_tokenizer_contract["full_vocab_size"], ), ( "student.tokenizer_fingerprint", student_meta.get("tokenizer_fingerprint"), student_tokenizer_contract["fingerprint"], ), ] warn_only_fields = { "teacher.tokenizer_fingerprint", "student.tokenizer_fingerprint", } for field_name, found, expected in expected_pairs: if found != expected: if field_name in warn_only_fields: log.warning( f" Provenance WARNING (non-fatal): {field_name} is {found!r}, " f"expected {expected!r}. This is likely due to a transformers " f"library version change. Continuing because vocab sizes match." ) else: log.error( f"PROVENANCE MISMATCH: {field_name} is {found!r}, expected {expected!r}." ) log.error("Regenerate compatible teacher-logit shards.") raise SystemExit(1) provenance_data_path = prov.get("data_path") current_data_path = str(Path(data_path).resolve()) if provenance_data_path != current_data_path: log.warning( " Provenance data_path differs because logits were likely generated on another machine: " f"{provenance_data_path!r} vs {current_data_path!r}. " "Continuing because data_sha256 matches." ) shard_schema = prov.get("shard_schema") expected_shard_schema = build_shard_schema() if shard_schema != expected_shard_schema: log.error( f"PROVENANCE MISMATCH: shard_schema is {shard_schema!r}, " f"expected {expected_shard_schema!r}." ) log.error("Regenerate compatible teacher-logit shards.") raise SystemExit(1) log.info(" Provenance: PASS (teacher shards match the current config and dataset)")