Voice / vlib /cache.py
Wiself's picture
Update voice.py, vlib, tests, bundles
4a1430c
Raw
History Blame Contribute Delete
15.3 kB
"""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}")