Voice / vlib /fetch.py
Wiself's picture
Update voice.py, vlib, tests, bundles
fcff745
Raw History Blame Contribute Delete
8.15 kB
"""Fetching tensors from sources into local bytes."""
from pathlib import Path
import json
import os
from vlib import ctx
from vlib.ui import _fail, _now_iso, _step, _warn
from vlib.net import _cleanup_tmp, _safe_id, download_range
from vlib.tensors import _is_ggml_quant, _is_quant_fetched, _write_safetensors_streaming, _write_single_tensor_gguf, encode_from_f32, write_safetensors
from vlib.sources import _hf_config, _public_source, detect_output_tensor, largest_2d
def _source_arch(source):
if source.kind == "safetensors":
cfg = _hf_config(source.repo) if source.repo else None
if cfg:
arch = (cfg.get("architectures") or [None])[0] or cfg.get("model_type")
else:
arch = None
return arch, cfg
source._load()
if getattr(source, "_reader", None) is not None:
for f in source._reader.fields.values():
if f.name == "general.architecture":
try:
v = f.parts[f.data[0]]
return (v.decode() if isinstance(v, bytes) else str(v)), None
except Exception:
return None, None
return (source._remote[0] if source._remote else None), None
def _resolve_output_name(source, config, names):
if source.kind == "gguf":
for n in ("output.weight", "token_embd.weight"):
if n in names:
return n
n = detect_output_tensor(names, config)
if n:
return n
else:
n = detect_output_tensor(names, config)
if n:
return n
shapes = {n: source.ref(n).shape for n in names}
return largest_2d(names, shapes)
def _stream_safetensors_tensor(source, name, dest_path):
"""Stream a safetensors tensor's raw bytes to dest_path. Returns byte count."""
if source.path:
hs, tensors = source._local()
begin, end = tensors[name]["data_offsets"]
with open(source.path, "rb") as src, open(dest_path, "wb") as out:
src.seek(8 + hs + begin)
remaining = end - begin
while remaining:
b = src.read(min(8 << 20, remaining))
if not b:
break
out.write(b)
remaining -= len(b)
return end - begin
url, hs, tensors = source._shard_info(name)
begin, end = tensors[name]["data_offsets"]
import shutil
combined = download_range(url, 8 + hs + begin, 8 + hs + end - 1, dest_path.parent, None, label=name)
shutil.move(str(combined), str(dest_path))
return end - begin
def _fetch_tensor(source, name, tmp_dir):
"""Return (kind, raw_path, data, dtype, shape) for one tensor.
kind == 'file' → raw bytes at raw_path; kind == 'bytes' → data holds bytes.
GGML-quantized tensors keep their exact raw blocks (no requant)."""
ref = source.ref(name)
tmp_dir.mkdir(parents=True, exist_ok=True)
if source.kind == "safetensors":
raw_path = tmp_dir / "raw.bin"
n = _stream_safetensors_tensor(source, name, raw_path)
return ("file", raw_path, n, ref.dtype, ref.shape)
if _is_ggml_quant(ref.dtype):
# passthrough: exact quant blocks, stored in a single-tensor GGUF artifact
raw_path = tmp_dir / "raw.bin"
n = _stream_gguf_tensor(source, name, raw_path)
return ("file", raw_path, n, ref.dtype, ref.shape)
f32 = source.read_f32(name)
# Preserve original precision family: BF16 stays BF16 (not F16 downconvert),
# F32 stays F32, F16 stays F16. BF16->F32->BF16 round-trip is exact.
_rd = ref.dtype.upper()
if _rd in ("F32", "FP32"):
dtype = "F32"
elif _rd in ("BF16", "BF16E8M"):
dtype = "BF16"
else:
dtype = "F16"
return ("bytes", None, encode_from_f32(f32, dtype), dtype, ref.shape)
def _stream_gguf_tensor(source, name, dest_path):
"""Stream a GGUF tensor's exact raw blocks to dest_path. Returns byte count."""
import shutil
import numpy as np
if source.path:
source._load()
for t in source._reader.tensors:
if t.name == name:
with open(dest_path, "wb") as out:
raw = t.data.tobytes() if t.data.dtype == np.uint8 else memoryview(t.data).cast("B").tobytes()
out.write(raw)
return t.n_bytes
raise KeyError(name)
url, parsed = source._owner(name)
dims, gtype, data_offset, n_bytes = parsed[2][name]
start = parsed[1] + data_offset
combined = download_range(url, start, start + n_bytes - 1, dest_path.parent, None, label=name)
shutil.move(str(combined), str(dest_path))
return n_bytes
def _write_fetched(fetched, tensor_name, out_path):
"""Write a fetched tensor to a standalone artifact (gguf when quantized)."""
kind, raw_path, data, dtype, shape = fetched
if _is_quant_fetched(fetched):
if not str(out_path).endswith(".gguf"):
_fixed = str(out_path)[:-len(".safetensors")] + ".gguf" if str(out_path).endswith(".safetensors") else str(out_path) + ".gguf"
_warn(f" Quant artifact needs .gguf — writing {_fixed} instead of the requested extension.")
out_path = _fixed
_write_single_tensor_gguf(tensor_name, raw_path, raw_path.stat().st_size, dtype, shape, str(out_path))
return Path(out_path)
tmp = str(out_path) + ".tmp"
if kind == "file":
_write_safetensors_streaming(tensor_name, str(raw_path), raw_path.stat().st_size, dtype, shape, tmp)
else:
write_safetensors({tensor_name: (data, dtype, shape)}, tmp)
os.replace(tmp, str(out_path))
return Path(out_path)
def _write_voicepack_json(out_path, source_id, pack):
"""Provenance sidecar for a voicepack: source + per-tensor names/shapes/dtypes.
Written next to the pack file (same stem, .json suffix)."""
meta = {
"source": _public_source(source_id),
"tensors": [{"name": n, "shape": [int(x) for x in v[2]], "dtype": v[1]}
for n, v in pack.items()],
"downloaded_at": _now_iso(),
}
jpath = Path(str(out_path)).with_suffix(".json")
jtmp = jpath.parent / (jpath.name + ".tmp")
jtmp.write_text(json.dumps(meta, indent=2) + "\n")
jtmp.replace(jpath)
return jpath
def _fetch_pack(source, source_id, resolved):
"""Fetch each resolved tensor's raw bytes. Returns pack {name: (bytes, dtype, shape)}.
Shared by get --multi and pack: always safetensors (GGUF quants dequant once with a warning)."""
pack = {}
for tname in resolved:
ref = source.ref(tname)
_step(f"Extracting {tname} ({ref.dtype} {tuple(ref.shape)})…")
tmp_dir = ctx.VOICES_DIR / ".parts" / _safe_id(source_id, tname)
try:
fetched = _fetch_tensor(source, tname, tmp_dir)
except Exception as e:
_fail(f" ✗ Could not extract tensor {tname}: {e}")
kind, raw_path, data, dtype, shape = fetched
try:
if kind == "file":
# safetensors raw file: read bytes directly (preserves BF16/F16/F32).
with open(str(raw_path), "rb") as f:
b = f.read()
pack[tname] = (b, dtype, tuple(shape))
else:
# bytes already encoded (F32/F16/BF16 preserved).
pack[tname] = (data, dtype, tuple(shape))
# If fetched was GGUF quant passthrough, _fetch_tensor returns
# ("file", raw, n, qtype, shape) with quant dtype — dequant once for safetensors pack.
if _is_ggml_quant(dtype):
_warn(f" {tname} is {dtype} quant — dequanting once to F32 for voicepack.safetensors (unavoidable, safetensors has no quant).")
try:
f32 = source.read_f32(tname)
pack[tname] = (f32.astype("float32").tobytes(), "F32", tuple(shape))
except Exception as e:
_fail(f" ✗ Could not dequant {tname} for pack: {e}")
finally:
_cleanup_tmp(tmp_dir)
if not pack:
_fail(" ✗ Nothing to pack.")
return pack