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