File size: 9,199 Bytes
8fc0f7f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
"""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")