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}")
|