File size: 8,563 Bytes
4335e83
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
"""Phase B: dump MTP training triples from the trunk over the corpus.

For each corpus record, renders the chat template over (messages + completion),
runs the quantized trunk once, and stores per position t:

    hidden[t]        post-norm_f trunk hidden (bf16)   -- the MTP head input
    token[t+1]       next token id                     -- fused with hidden[t]
    token[t+2]       hard target
    topk_logits[t+1] trunk's top-K logits at t+1       -- soft target

The soft target is the point: the trunk's distribution at position t+1 IS
P(t+2 | prefix, t+1) — exactly what the head should predict given
(hidden[t], token[t+1]). Training with KL to this aligns the head to the
trunk's conditionals regardless of who authored the text; the corpus only
controls state coverage. Loss = CE(hard) + lambda * KL(soft), lambda ~1.

Storage: one .npz per corpus shard chunk in triples/:
    hiddens  (N, H) bf16 saved as uint16 view
    tokens   (N+2,) uint32   (positions align: hidden[i] pairs tokens[i+1], tokens[i+2])
    topk_ids (N, K) uint32; topk_lp (N, K) float16   (K=64, log-softmaxed)
    doc_bounds: (n_docs, 2) start/end into N (docs are independent sequences)

Only completion-region positions are dumped (prompt positions carry
off-distribution states for the head's serving regime, where it drafts during
generation). A position budget per doc caps degenerate long docs.

MEMORY: loads the full quantized model (~81 GB) via mlx_lm. Run while the
omlx server is idle or stopped if headroom is tight; peak ~90 GB + activations.

Usage: python3 dump_triples.py [--corpus-glob 'corpus/*.jsonl'] [--out triples]
                               [--topk 64] [--max-docs N] [--resume]

--resume continues after whatever triples-*.npz already exist (numbering and
doc position carry on from them). Without it, a non-empty --out is an error
rather than a silent overwrite.
"""

import argparse
import glob as globmod
import json
import time
from pathlib import Path

import mlx.core as mx
import numpy as np

MODEL_DIR = "/Users/david/AI/kenosistron3-oq5e-mtp"
ROOT = Path(__file__).parent
CHUNK_TOKENS = 4096          # prefill chunk (fits the SSD kernel fast path)
FLUSH_POSITIONS = 200_000    # positions per output .npz (~1.6 GB hiddens)


def load_model():
    # omlx's oq loader handles the oQ5e quantization; mlx_lm.load handles the
    # rest of the architecture. Try mlx_lm first; fall back to omlx.
    try:
        from mlx_lm import load

        model, tokenizer = load(MODEL_DIR)
        return model, tokenizer
    except Exception as e:
        raise SystemExit(
            f"mlx_lm.load failed ({e}); if this is an oQ-format issue, run "
            "through omlx's loader instead (omlx.utils.model_loading)."
        )


def render(tokenizer, messages, completion):
    """Token ids for the full exchange + index where the completion starts."""
    prompt_ids = tokenizer.apply_chat_template(
        messages, add_generation_prompt=True, tokenize=True
    )
    comp_ids = tokenizer.encode(completion, add_special_tokens=False)
    return list(prompt_ids) + list(comp_ids), len(prompt_ids)


def trunk_hiddens(model, ids):
    """Post-norm_f hiddens + logits for a full sequence, chunked prefill."""
    cache = model.make_cache()
    hiddens = []
    logits = []
    x = mx.array(ids)[None]
    for s0 in range(0, x.shape[1], CHUNK_TOKENS):
        chunk = x[:, s0 : s0 + CHUNK_TOKENS]
        lg, hd = model(chunk, cache=cache, return_hidden=True)
        mx.eval(lg, hd)
        hiddens.append(hd[0])
        logits.append(lg[0])
    return mx.concatenate(hiddens), mx.concatenate(logits)


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--corpus-glob", default=str(ROOT / "corpus" / "*.jsonl"))
    ap.add_argument("--out", default=str(ROOT / "triples"))
    ap.add_argument("--topk", type=int, default=64)
    ap.add_argument("--max-docs", type=int, default=None)
    ap.add_argument("--max-positions-per-doc", type=int, default=1600)
    ap.add_argument(
        "--resume",
        action="store_true",
        help="continue after the existing triples-*.npz instead of overwriting",
    )
    args = ap.parse_args()
    out = Path(args.out)
    out.mkdir(exist_ok=True)

    # The MTP-side patches must be active so return_hidden exists.
    from omlx.patches.mlx_lm_mtp import apply_mlx_lm_mtp_patch

    apply_mlx_lm_mtp_patch()

    model, tokenizer = load_model()
    docs = []
    for f in sorted(globmod.glob(args.corpus_glob)):
        for line in open(f):
            docs.append(json.loads(line))
    if args.max_docs:
        docs = docs[: args.max_docs]
    print(f"{len(docs)} corpus docs")

    existing = sorted(out.glob("triples-*.npz"))
    start_di, docs_done, shard_start = 0, 0, 0
    if existing:
        if not args.resume:
            raise SystemExit(
                f"{len(existing)} shard(s) already in {out}; pass --resume to "
                "continue after them (re-running without it overwrites from 0000)."
            )
        shard_start = len(existing)
        for f in existing:
            with np.load(f) as d:
                docs_done += d["doc_bounds"].shape[0]
        # Replay the filters below to find where those docs ended. They depend
        # only on render() and lengths -- not the trunk -- so the fast-forward
        # is exact and costs no model forwards.
        seen, start_di = 0, len(docs)
        for di, doc in enumerate(docs):
            if seen >= docs_done:
                start_di = di
                break
            ids, comp_start = render(tokenizer, doc["messages"], doc["completion"])
            if len(ids) - comp_start < 8:
                continue
            lo = max(comp_start - 1, 0)
            hi = min(len(ids) - 2, lo + args.max_positions_per_doc)
            if hi <= lo:
                continue
            seen += 1
        if seen < docs_done:
            raise SystemExit(
                f"corpus yields only {seen} usable docs but shards hold "
                f"{docs_done}; corpus-glob likely differs from the first run."
            )
        print(
            f"resume: {docs_done} docs in {shard_start} shard(s), "
            f"restarting at corpus index {start_di}",
            flush=True,
        )

    buf_h, buf_tok, buf_ids, buf_lp, bounds = [], [], [], [], []
    n_pos, shard_i, t0 = 0, shard_start, time.time()

    def flush():
        nonlocal buf_h, buf_tok, buf_ids, buf_lp, bounds, shard_i
        if not buf_h:
            return
        np.savez(
            out / f"triples-{shard_i:04d}.npz",
            hiddens=np.concatenate(buf_h),
            tokens=np.concatenate(buf_tok),
            topk_ids=np.concatenate(buf_ids),
            topk_lp=np.concatenate(buf_lp),
            doc_bounds=np.array(bounds, dtype=np.int64),
        )
        print(f"  wrote triples-{shard_i:04d}.npz")
        shard_i += 1
        buf_h, buf_tok, buf_ids, buf_lp, bounds = [], [], [], [], []

    written = docs_done
    for di, doc in enumerate(docs):
        if di < start_di:
            continue
        ids, comp_start = render(tokenizer, doc["messages"], doc["completion"])
        if len(ids) - comp_start < 8:
            continue
        h, lg = trunk_hiddens(model, ids)
        # positions t in the completion region with t+2 in range
        lo = max(comp_start - 1, 0)
        hi = min(len(ids) - 2, lo + args.max_positions_per_doc)
        if hi <= lo:
            continue
        hs = h[lo:hi]                                   # (n, H)
        lg1 = lg[lo + 1 : hi + 1].astype(mx.float32)    # trunk logits at t+1
        lp1 = lg1 - mx.logsumexp(lg1, axis=-1, keepdims=True)
        kidx = mx.argpartition(-lp1, kth=args.topk - 1, axis=-1)[:, : args.topk]
        klp = mx.take_along_axis(lp1, kidx, axis=-1)
        mx.eval(hs, kidx, klp)
        start = sum(x.shape[0] for x in buf_h)
        buf_h.append(np.array(hs.astype(mx.bfloat16).view(mx.uint16)))
        buf_tok.append(np.array(ids[lo : hi + 2], dtype=np.uint32))
        buf_ids.append(np.array(kidx.astype(mx.uint32)))
        buf_lp.append(np.array(klp.astype(mx.float16)))
        bounds.append((start, start + (hi - lo)))
        n_pos += hi - lo
        written += 1
        if n_pos >= FLUSH_POSITIONS:
            flush()
            n_pos = 0
        if written % 50 == 0:
            dt = time.time() - t0
            print(f"[{dt/60:5.1f}m] {written}/{len(docs)} docs", flush=True)
    flush()
    print(f"DONE: {written} docs in {(time.time()-t0)/60:.1f} min -> {out}")


if __name__ == "__main__":
    main()