File size: 15,266 Bytes
90f10d2
 
 
 
 
 
 
91d026f
90f10d2
 
4a1430c
90f10d2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4a1430c
 
 
 
 
 
 
 
 
91d026f
 
 
 
 
 
 
 
4a1430c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
91d026f
4a1430c
91d026f
 
90f10d2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fcff745
4a1430c
 
 
 
fcff745
90f10d2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
"""Base-model tensor cache."""

from pathlib import Path
import json
import os
import struct
from vlib import ctx
from vlib.ui import _fail, _now_iso, _say, _step, _warn
from vlib.net import _cleanup_tmp, _safe_id, http_get_json
from vlib.tensors import _is_ggml_quant, _write_safetensors_streaming, read_safetensors, write_safetensors
from vlib.sources import _base_paths, open_source, role_of, is_tied
from vlib.fetch import _fetch_tensor, _resolve_output_name, _source_arch
from vlib.registry import _assert_not_symlink


def _base_rev(base_id, base_file):
    """Pinned commit sha for a base cache: sidecar .rev wins, else API HEAD (then pin)."""
    rev_path = base_file.with_suffix(".rev")
    try:
        if rev_path.exists():
            rev = rev_path.read_text().strip()
            if rev:
                return rev
    except Exception:
        pass
    rev = None
    try:
        rev = http_get_json(f"https://huggingface.co/api/models/{base_id}", timeout=20).get("sha")
    except Exception:
        rev = None
    if rev:
        try:
            base_file.parent.mkdir(parents=True, exist_ok=True)
            tmp = rev_path.with_suffix(".rev.tmp")
            tmp.write_text(rev + "\n")
            tmp.replace(rev_path)
        except Exception:
            pass
        return rev
    return "main"


def _cache_validate(base_file, names):
    """Names neither in the header nor fully covered by file bytes. Empty = trusted."""
    try:
        with open(base_file, "rb") as f:
            hs = struct.unpack("<Q", f.read(8))[0]
            hdr = json.loads(f.read(hs))
        size = os.path.getsize(base_file)
    except Exception:
        return list(names)
    missing = []
    for n in names:
        info = hdr.get(n)
        if not isinstance(info, dict):
            missing.append(n)
            continue
        try:
            off = info["data_offsets"]
            if not (isinstance(off, (list, tuple)) and len(off) == 2
                    and 0 <= int(off[0]) <= int(off[1]) and 8 + hs + int(off[1]) <= size):
                missing.append(n)
        except Exception:
            missing.append(n)
    return missing


def _cache_topup_write(base_file, staged):
    """Append staged tensors to a base cache + atomically rewrite its header.
    staged: [(name, blob|None, dtype, shape, raw_path|None)] β€” blob XOR raw_path.
    Offsets are data-relative so existing entries never shift. Returns new header.
    Raises OSError on failure; base_file is only ever swapped in whole (tmp+rename)."""
    _assert_not_symlink(base_file)
    import shutil
    # append + atomic header rewrite (offsets are data-relative: existing entries don't shift)
    with open(base_file, "rb") as f:
        old_hs = struct.unpack("<Q", f.read(8))[0]
        old_hdr = json.loads(f.read(old_hs))
    old_size = os.path.getsize(base_file)
    old_data_len = old_size - 8 - old_hs
    new_hdr = {k: v for k, v in old_hdr.items()}
    off = old_data_len
    for m, blob, dtype, shape, _rp in staged:
        nbytes = len(blob) if blob is not None else Path(_rp).stat().st_size
        new_hdr[m] = {"dtype": dtype, "shape": list(shape), "data_offsets": [off, off + nbytes]}
        off += nbytes
    new_hj = json.dumps(new_hdr).encode("utf-8")
    tmp_new = base_file.with_suffix(".cache.tmp")
    with open(tmp_new, "wb") as out:
        out.write(struct.pack("<Q", len(new_hj)))
        out.write(new_hj)
        with open(base_file, "rb") as f:
            f.seek(8 + old_hs)
            shutil.copyfileobj(f, out, 1 << 20)
        for _m, blob, _dt, _sh, _rp in staged:
            if blob is not None:
                out.write(blob)
            else:
                with open(_rp, "rb") as f:
                    shutil.copyfileobj(f, out, 1 << 20)
        out.flush()
        try:
            os.fsync(out.fileno())
        except Exception:
            pass
    try:
        try:
            os.chmod(tmp_new, 0o600)
        except Exception:
            pass
        tmp_new.replace(base_file)
    except Exception as e:
        try:
            tmp_new.unlink(missing_ok=True)
        except Exception:
            pass
        raise OSError(f"Could not update base cache: {e}")
    return new_hdr


def _resolve_head_want(base_id, rev, want):
    """Reroute wanted head/embed tensors onto the base's same-role names.
    Covers cross-format aliases (voice `output.weight` vs base `lm_head`)
    and tied bases (voice head vs base embed β€” the delta pairs them later).
    Returns the want list, order-preserved and deduplicated. Anything
    unresolvable (offline, exact names present, no same-role counterpart)
    returns want unchanged and downstream fails exactly as before."""
    heads = [n for n in want if role_of(n) in ("head", "embed")]
    if not heads:
        return want
    try:
        src = open_source(base_id, rev=rev)
        names = [n for n in src.names() if n != "delta.voice.marker"]
    except SystemExit:
        raise
    except Exception:
        return want
    remap = {}
    for h in heads:
        if h in names:
            continue
        cands = [n for n in names if role_of(n) == role_of(h)]
        tied = False
        if not cands and role_of(h) == "head" and is_tied(names):
            cands = [n for n in names if role_of(n) == "embed"]
            tied = bool(cands)
        if len(cands) != 1:
            return want
        remap[h] = cands[0]
        if tied:
            _step("Base model ties its head β€” fetching its embedding for the match.")
    if not remap:
        return want
    return list(dict.fromkeys(remap.get(n, n) for n in want))


def _ensure_base_tensors(base_id, names, args=None):
    """Base cache with per-tensor trust: top-ups what's missing, self-heals partial caches.
    All fetches pin one revision so a repo update mid-cache can't mix commits. Returns base_file."""
    tmp = _base_paths(base_id)
    if tmp is None:
        _fail(f"  βœ— Invalid base '{base_id}'")
    base_file, base_json = tmp
    want = [n for n in names if n != "delta.voice.marker"]
    if not want:
        _fail("  βœ— No tensors requested from base.")
    import fcntl
    base_file.parent.mkdir(parents=True, exist_ok=True)
    lock_path = base_file.with_suffix(".lock")
    try:
        lock_fh = open(lock_path, "w")
    except Exception:
        lock_fh = None
    try:
        if lock_fh is not None:
            try:
                fcntl.flock(lock_fh, fcntl.LOCK_EX)
            except Exception:
                pass
        return _ensure_base_tensors_locked(base_id, base_file, base_json, want)
    finally:
        if lock_fh is not None:
            try:
                fcntl.flock(lock_fh, fcntl.LOCK_UN)
            except Exception:
                pass
            try:
                lock_fh.close()
            except Exception:
                pass


def _ensure_base_tensors_locked(base_id, base_file, base_json, want):
    import shutil
    rev_path = base_file.with_suffix(".rev")
    legacy = base_file.exists() and not rev_path.exists()
    if not base_file.exists():
        _cache_base(base_id, base_file, base_json)
    elif legacy:
        # Pre-rev cache of unknown lineage: re-fetch one held tensor at HEAD and
        # compare bytes. Match -> same lineage, pin HEAD. Differ -> wipe, start over.
        # Probe failure (offline?) keeps the cache with a loud warning, never wipes.
        try:
            with open(base_file, "rb") as f:
                hs = struct.unpack("<Q", f.read(8))[0]
                hdr = json.loads(f.read(hs))
            held = [k for k in hdr.keys() if k != "__metadata__"]
        except Exception:
            held = []
        verified, same = False, False
        if held:
            probe = held[0]
            probe_dir = ctx.VOICES_DIR / ".parts" / _safe_id("base-probe", base_id)
            try:
                src = open_source(base_id)
                pf = _fetch_tensor(src, probe, probe_dir)
                fresh = Path(pf[1]).read_bytes() if pf[0] == "file" else bytes(pf[2])
                with open(base_file, "rb") as f:
                    f.seek(8 + hs + hdr[probe]["data_offsets"][0])
                    have = f.read(hdr[probe]["data_offsets"][1] - hdr[probe]["data_offsets"][0])
                verified, same = True, (fresh == have)
            except SystemExit:
                raise
            except Exception:
                verified, same = False, False
            finally:
                _cleanup_tmp(probe_dir)
        if verified and not same:
            _warn("  Base cache predates revision pinning and no longer matches HEAD β€” re-fetching.")
            try:
                base_file.unlink(missing_ok=True)
            except Exception:
                pass
            _cache_base(base_id, base_file, base_json)
        elif not verified:
            _warn("  Base cache lineage unverified (offline?) β€” proceeding, mixed revisions possible.")
    rev = _base_rev(base_id, base_file)
    missing = _cache_validate(base_file, want)
    if missing:
        # Reroute before fetching: same-role base names serve cross-format
        # heads/embeds (voice head vs tied base serves the embed instead).
        # Probed only on a miss, so warm caches never touch network.
        want = _resolve_head_want(base_id, rev, want)
        missing = _cache_validate(base_file, want)
    if not missing:
        return base_file
    _step(f"Topping up base {base_id} ({len(missing)} tensor(s) missing)…")
    try:
        src = open_source(base_id, rev=rev)
    except SystemExit:
        raise
    except Exception as e:
        _fail(f"  βœ— Could not reach base {base_id} for top-up: {e}")
    tmp_dir = ctx.VOICES_DIR / ".parts" / _safe_id("base", base_id)
    staged = []
    try:
        for m in missing:
            try:
                fetched = _fetch_tensor(src, m, tmp_dir)
            except (KeyError, ValueError) as e:
                _fail(f"  βœ— Base {base_id} has no tensor '{m}'. Delta needs same names on both sides.")
            kind, raw_path, data, dtype, shape = fetched
            if _is_ggml_quant(dtype):
                # quant blocks can't live in a .safetensors cache: dequant once to F32
                f32 = src.read_f32(m)
                staged.append((m, f32.astype("float32").tobytes(), "F32", tuple(int(x) for x in f32.shape), None))
            elif kind == "file":
                staged.append((m, None, dtype, tuple(int(x) for x in shape), raw_path))
            else:
                staged.append((m, bytes(data), dtype, tuple(int(x) for x in shape), None))
    except SystemExit:
        _cleanup_tmp(tmp_dir)
        raise
    try:
        new_hdr = _cache_topup_write(base_file, staged)
    except OSError as e:
        _cleanup_tmp(tmp_dir)
        _fail(f"  βœ— {e}")
    _cleanup_tmp(tmp_dir)
    # refresh base json with the full tensor list
    try:
        meta = {"source_hf_model": base_id, "revision": rev, "downloaded_at": _now_iso(),
                "tensors": [{"name": k, "shape": list(v["shape"]), "dtype": v["dtype"]}
                            for k, v in new_hdr.items() if k != "__metadata__"]}
        jtmp = base_json.with_suffix(".tmp")
        fd = os.open(str(jtmp), os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600)
        try:
            os.write(fd, json.dumps(meta, indent=2).encode("utf-8") + b"\n")
            try:
                os.fsync(fd)
            except Exception:
                pass
        finally:
            os.close(fd)
        try:
            jtmp.replace(base_json)
        except FileExistsError:
            jtmp.unlink(missing_ok=True)
    except Exception:
        pass
    return base_file


def _cache_base(base_id, base_file, base_json):
    _step(f"Caching base {base_id}…")
    src = open_source(base_id)
    names = src.names()
    bname = _resolve_output_name(src, None, names)
    if bname is None:
        _fail(f"  βœ— Could not find the output tensor in base {base_id}.")
    ref = src.ref(bname)
    tmp_dir = ctx.VOICES_DIR / ".parts" / _safe_id("base", base_id)
    fetched = _fetch_tensor(src, bname, tmp_dir)
    # size now: _cleanup_tmp below deletes raw.bin before the metadata write
    fetched_bytes = fetched[1].stat().st_size if fetched[0] == "file" else len(fetched[2])
    _assert_not_symlink(base_file.parent if base_file.parent.exists() else base_file)
    base_file.parent.mkdir(parents=True, exist_ok=True)
    try:
        os.chmod(base_file.parent, 0o700)
    except Exception:
        pass
    # avoid TOCTOU race if two processes cache same base
    if base_file.exists():
        _cleanup_tmp(tmp_dir)
        return
    tmp_st = base_file.parent / "base.tmp"
    if fetched[0] == "file":
        _write_safetensors_streaming(bname, str(fetched[1]), fetched[1].stat().st_size, fetched[3], fetched[4], str(tmp_st))
    else:
        write_safetensors({bname: (fetched[2], fetched[3], fetched[4])}, str(tmp_st))
    try:
        tmp_st.replace(base_file)
    except FileExistsError:
        tmp_st.unlink(missing_ok=True)
    _cleanup_tmp(tmp_dir)
    # atomic 0600 json
    jtmp = base_json.with_suffix(".tmp")
    fd = os.open(str(jtmp), os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600)
    try:
        os.write(fd, json.dumps({
            "source_hf_model": base_id, "tensor_name": bname, "dtype": fetched[3],
            "shape": list(fetched[4]), "bytes": fetched_bytes,
            "tensors": [{"name": bname, "shape": list(fetched[4]), "dtype": fetched[3]}],
            "downloaded_at": _now_iso(),
        }, indent=2).encode("utf-8") + b"\n")
        try:
            os.fsync(fd)
        except Exception:
            pass
    finally:
        os.close(fd)
    try:
        jtmp.replace(base_json)
    except FileExistsError:
        jtmp.unlink(missing_ok=True)


def _load_safetensors_or_gguf_f32(path):
    # Returns (tensors_dict, dtype_map) for delta math. Safetensors preferred (no dequant).
    # GGUF quant falls back to full F32 dequant once (unavoidable, warn; efficient for <50MB tests).
    try:
        hdr, tens = read_safetensors(str(path))
        return hdr, tens, False
    except Exception:
        pass
    try:
        src = open_source(str(path))
        names = [n for n in src.names() if n != "delta.voice.marker"]
        if not names:
            raise ValueError("no tensors")
        # Pick output head or first for delta math (single-tensor delta path).
        tname = None
        try:
            arch, cfg = _source_arch(src)
            tname = _resolve_output_name(src, cfg, names)
        except Exception:
            tname = None
        if tname is None or tname not in names:
            tname = names[0]
        ref = src.ref(tname)
        _warn(f"  {Path(path).name} is {ref.dtype} ({src.kind}) β€” dequanting once to F32 for delta math (unavoidable).")
        arr = src.read_f32(tname)
        # Build synthetic safetensors-like dict for chunked math below.
        return {"__gguf_f32__": {"dtype": "F32", "shape": list(arr.shape), "_arr": arr}, tname: {"dtype": "F32", "shape": list(arr.shape), "_arr": arr}}, {tname: {"dtype": "F32", "shape": list(arr.shape)}}, True
    except Exception as e:
        raise ValueError(f"{path} is not a valid .safetensors or .gguf for delta: {e}")