safffrron's picture
Upload code.py with huggingface_hub
8fc0f7f verified
Raw
History Blame Contribute Delete
9.2 kB
"""Week-2 20% mixed-bit base plus grouped-JSD residual entry points."""
from __future__ import annotations
import importlib.util
import json
import os
import shutil
import sys
import sysconfig
import tempfile
from pathlib import Path
# The course requires this filename, which otherwise shadows Python's stdlib
# ``code`` module while torch imports pdb. Export the stdlib API first.
if __name__ == "code":
_stdlib_path = Path(sysconfig.get_path("stdlib")) / "code.py"
_spec = importlib.util.spec_from_file_location("_cs6013_stdlib_code", _stdlib_path)
if _spec is None or _spec.loader is None:
raise ImportError(f"could not load stdlib code module: {_stdlib_path}")
_stdlib = importlib.util.module_from_spec(_spec)
_spec.loader.exec_module(_stdlib)
for _name in ("InteractiveInterpreter", "InteractiveConsole", "interact", "compile_command"):
globals()[_name] = getattr(_stdlib, _name)
import torch
LOCAL_SRC = Path(__file__).resolve().parent / "src"
if LOCAL_SRC.is_dir() and str(LOCAL_SRC) not in sys.path:
sys.path.insert(0, str(LOCAL_SRC))
from eaimath.artifact import ( # noqa: E402
dequantize_gptq_model,
pack_model_state,
read_trace_corpus,
restore_artifact,
save_artifact,
)
from eaimath.buckets import BY_NAME # noqa: E402
from eaimath.embedding_predictor import fit_token_predictor # noqa: E402
from eaimath.model import load_model, load_tokenizer, save_checkpoint # noqa: E402
from eaimath.pack import state_dict_bytes # noqa: E402
from eaimath.peft_compat import prepare_dense_lora_dispatch # noqa: E402
from eaimath.vocab import build_keep_set, token_frequencies # noqa: E402
SUBMISSION_HF_REPO = "safffrron/25M2111-Week02-Track2-20-Submission01"
BASE_FILENAME = "week02_20_base.ptz"
KEEP_FILENAME = "week02_20_keep_ids.json"
ADAPTER_DIRNAME = "adapter"
EXPECTED_BITS = {"mlp": 3, "linear_attn": 4, "full_attn": 8, "embed": 8}
def _reproduction_inputs() -> tuple[Path, Path, Path, dict]:
source = Path(os.environ.get("EAIMATH_B20_GPTQ_SOURCE", "work/gptq_m3l4a8e8"))
corpus = Path(os.environ.get("EAIMATH_VOCAB_CORPUS", "data/traces.jsonl"))
adapter = Path(os.environ.get("EAIMATH_WEEK2_ADAPTER_SOURCE", "work/b_postq_jsd_r8"))
required = [source / "experiment_config.json", corpus, adapter / "adapter_model.safetensors", adapter / "adapter_config.json"]
missing = [str(path) for path in required if not path.is_file()]
if missing:
raise FileNotFoundError(
"Reproduction inputs are missing. Follow README.md's rebuild section or set "
"EAIMATH_B20_GPTQ_SOURCE, EAIMATH_VOCAB_CORPUS, and "
f"EAIMATH_WEEK2_ADAPTER_SOURCE. Missing: {missing}"
)
config = json.loads((source / "experiment_config.json").read_text())
if config.get("bits") != EXPECTED_BITS:
raise ValueError(f"expected m3l4a8e8, found bits={config.get('bits')}")
if set(config.get("gptq_components", [])) != {"mlp", "linear_attn"}:
raise ValueError("GPTQ source must calibrate both MLP and linear attention")
if config.get("activation_order") != "gar":
raise ValueError("GPTQ source must use GAR activation order")
adapter_config = json.loads((adapter / "adapter_config.json").read_text())
if int(adapter_config.get("r", -1)) != 8:
raise ValueError("the submitted residual must have rank 8")
return source, corpus, adapter, config
def _locate_payload(checkpoint_path: str) -> tuple[Path, Path, Path]:
supplied = Path(checkpoint_path).expanduser().resolve()
root = supplied if supplied.is_dir() else supplied.parent
base = root / BASE_FILENAME if supplied.is_dir() else supplied
keep = root / KEEP_FILENAME
adapter = root / ADAPTER_DIRNAME
if base.is_file() and keep.is_file() and (adapter / "adapter_model.safetensors").is_file():
return base, keep, adapter
from huggingface_hub import snapshot_download
snapshot = Path(
snapshot_download(
SUBMISSION_HF_REPO,
allow_patterns=[BASE_FILENAME, f"{BASE_FILENAME}.json", KEEP_FILENAME, "adapter/*"],
)
)
base, keep, adapter = snapshot / BASE_FILENAME, snapshot / KEEP_FILENAME, snapshot / ADAPTER_DIRNAME
missing = [str(path) for path in (base, keep, adapter / "adapter_model.safetensors", adapter / "adapter_config.json") if not path.is_file()]
if missing:
raise FileNotFoundError(f"incomplete Week-2 20% payload: {missing}")
return base, keep, adapter
def convert_from_hf_checkpoint(
model_name: str,
output_path: str,
sparsity: float | None = None,
) -> None:
"""Reproduce the packed base and copy the trained JSD residual beside it.
``sparsity`` is accepted only because the starter interface may pass it;
this method is dense mixed-bit quantization and does not use that value.
"""
_ = sparsity
source, corpus, adapter_source, source_config = _reproduction_inputs()
requested = Path(output_path)
if requested.suffix:
root, artifact = requested.parent, requested
else:
root, artifact = requested, requested / BASE_FILENAME
root.mkdir(parents=True, exist_ok=True)
if any((root / name).exists() for name in (KEEP_FILENAME, ADAPTER_DIRNAME)):
raise FileExistsError(f"refusing to overwrite an existing representation in {root}")
problems, completions, corpus_stats = read_trace_corpus([corpus])
tokenizer = load_tokenizer(str(source))
counts = token_frequencies(problems + completions, tokenizer)
priority = token_frequencies(problems, tokenizer)
keep_set = build_keep_set(counts, tokenizer, 30_000, priority_counts=priority)
loaded, model, replaced, exact_weights, capture_error = dequantize_gptq_model(source)
if capture_error != 0.0:
raise RuntimeError(f"GPTQ code capture error: {capture_error}")
payload, _ = pack_model_state(
model.state_dict(), keep_set["keep_ids"], group_size=128,
exact_entries=exact_weights, recipe_bits=EXPECTED_BITS, recipe_name="m3l4a8e8",
)
payload["source"] = {
"base_model": model_name,
"experiment_config": source_config,
"vocab_strategy": "problem-first",
"corpus_stats": corpus_stats,
"gptq_linears": replaced,
}
row_names = [name for name, entry in payload["state_dict"].items() if entry["kind"] in {"quant_rows", "raw_rows"}]
if len(row_names) != 1:
raise RuntimeError(f"expected one canonical tied embedding, found {row_names}")
predictor, predictor_report = fit_token_predictor(
model.state_dict()[row_names[0]].detach().cpu(), tokenizer,
sample_size=65_536, device=os.environ.get("EAIMATH_PREDICTOR_DEVICE", "cpu"),
)
payload["embedding_predictor"] = predictor
payload["report"]["embedding_predictor"] = predictor_report
payload["report"]["tensor_bytes"] += state_dict_bytes({"embedding_predictor": predictor["basis"]})
(root / KEEP_FILENAME).write_text(json.dumps(keep_set["keep_ids"]))
final = save_artifact(payload, artifact, lossless_codec="zlib", zlib_level=6)
shutil.copytree(adapter_source, root / ADAPTER_DIRNAME, ignore=shutil.ignore_patterns("__pycache__", "*.pyc"))
total = int(final["disk_bytes"]) + (root / KEEP_FILENAME).stat().st_size + sum(
path.stat().st_size for path in (root / ADAPTER_DIRNAME).rglob("*") if path.is_file()
)
ceiling = int(BY_NAME["B"].high_gb * 1e9)
if total >= ceiling:
raise RuntimeError(f"full representation exceeds the 20% ceiling: {total:,} >= {ceiling:,}")
del model, loaded
if torch.cuda.is_available():
torch.cuda.empty_cache()
def convert_to_hf_checkpoint(model_name: str, checkpoint_path: str, output_path: str) -> None:
"""Restore the packed base, safely merge the residual, and save BF16 HF."""
from peft import PeftModel
base, keep, adapter = _locate_payload(checkpoint_path)
output = Path(output_path)
if output.exists() and (not output.is_dir() or any(output.iterdir())):
raise FileExistsError(f"refusing to overwrite non-empty output: {output}")
output.parent.mkdir(parents=True, exist_ok=True)
device = os.environ.get("EAIMATH_RESTORE_DEVICE", "cpu")
with tempfile.TemporaryDirectory(prefix="week02-b20-", dir=output.parent) as temporary:
restored_base = Path(temporary) / "base"
base_report = restore_artifact(model_name, base, restored_base, embedding_fill="token")
model = load_model(str(restored_base), dtype="bfloat16", device=device, multimodal=True)
prepare_dense_lora_dispatch(model)
merged = PeftModel.from_pretrained(model, adapter).merge_and_unload(safe_merge=True)
merged.config.use_cache = True
copied = save_checkpoint(merged, output, source_model=model_name)
shutil.copy2(keep, output / "keep_ids.json")
report = {
"format": "eaimath-week02-b20-restored-v1",
"base_artifact": str(base),
"adapter": str(adapter),
"safe_merge": True,
"base_restore": base_report,
"auxiliary_files": copied,
}
(output / "submission_report.json").write_text(json.dumps(report, indent=2) + "\n")