"""Strict streaming loader for the versioned Prism custom-Q6 checkpoint.""" from __future__ import annotations import hashlib import json from pathlib import Path import time import torch from torch import nn from safetensors import safe_open from .quant import FORMAT, VERSION, PROFILE, GROUP_SIZE, Q6Linear, is_quantized_weight, packed_shape DTYPES = { "BF16": torch.bfloat16, "F16": torch.float16, "F32": torch.float32, "F64": torch.float64, "U8": torch.uint8, "I8": torch.int8, "I16": torch.int16, "I32": torch.int32, "I64": torch.int64, "BOOL": torch.bool, } def sha256_file(path: str | Path) -> str: digest = hashlib.sha256() with open(path, "rb") as stream: for block in iter(lambda: stream.read(8 * 1024 * 1024), b""): digest.update(block) return digest.hexdigest() def read_manifest(manifest_path: str | Path) -> tuple[Path, dict]: path = Path(manifest_path) if path.is_dir(): path = path / "manifest.json" if path.stat().st_size > 16 * 1024 * 1024: raise ValueError("Quantization manifest exceeds 16 MiB") data = json.loads(path.read_text()) if (data.get("format"), data.get("version"), data.get("profile"), data.get("group_size")) != ( FORMAT, VERSION, PROFILE, GROUP_SIZE, ): raise ValueError("Unsupported quantized checkpoint format/profile/version") for key in ("tensors", "shards", "quantized_modules", "source"): if not isinstance(data.get(key), dict): raise ValueError(f"Missing or invalid manifest {key}") for filename, record in data["shards"].items(): if Path(filename).name != filename or not filename.endswith(".safetensors"): raise ValueError("Unsafe shard filename") shard = path.parent / filename if shard.is_symlink() or not shard.is_file(): raise ValueError(f"Missing or symlinked shard: {filename}") if not isinstance(record.get("sha256"), str) or len(record["sha256"]) != 64: raise ValueError("Missing shard SHA-256") for name, record in data["tensors"].items(): if not isinstance(record, dict) or record.get("shard") not in data["shards"]: raise ValueError(f"Invalid tensor shard mapping: {name}") shape = record.get("shape") if not isinstance(shape, list) or any(type(n) is not int or n < 0 for n in shape): raise ValueError(f"Invalid tensor shape: {name}") if record.get("dtype") not in DTYPES: raise ValueError(f"Unsupported tensor dtype: {name}") return path, data def verify_shards(manifest_path: str | Path, manifest: dict | None = None) -> dict: path, parsed = read_manifest(manifest_path) if manifest is not None and manifest != parsed: raise ValueError("Manifest changed during loading") manifest = parsed for filename, record in manifest["shards"].items(): shard = path.parent / filename if shard.stat().st_size != record["bytes"] or sha256_file(shard) != record["sha256"]: raise ValueError(f"Shard length/hash mismatch: {filename}") expected = {k: v for k, v in manifest["tensors"].items() if v["shard"] == filename} with safe_open(str(shard), framework="pt", device="cpu") as handle: if set(handle.keys()) != set(expected): raise ValueError(f"Shard key mismatch: {filename}") for name, spec in expected.items(): view = handle.get_slice(name) if view.get_shape() != spec["shape"] or view.get_dtype() != spec["dtype"]: raise ValueError(f"Shard tensor shape/dtype mismatch: {name}") return manifest def load_quantized_model(model: nn.Module, manifest_path, device="cpu", strict=True): """Replace eligible Linears, assign one tensor at a time, return model/receipt. The caller constructs the assembled architecture without materialized weights. Nonpersistent buffers must already be real tensors; this function does not invent architecture-specific rotary frequencies or missing parameters. """ if not strict: raise ValueError("Partial quantized loading is unsupported; strict=True is required") started = time.monotonic() path, manifest = read_manifest(manifest_path) verify_shards(path, manifest) modules = dict(model.named_modules(remove_duplicate=False)) eligible = { name: module for name, module in modules.items() if isinstance(module, nn.Linear) and is_quantized_weight(name + ".weight", module.weight.shape) } definitions = manifest["quantized_modules"] if set(eligible) != set(definitions): raise ValueError(f"Quantized module mismatch: missing={sorted(set(eligible)-set(definitions))[:8]}, unexpected={sorted(set(definitions)-set(eligible))[:8]}") # Validate every target before replacing any of them. for name, original in eligible.items(): spec = definitions[name] if (spec.get("in_features"), spec.get("out_features"), spec.get("bias"), spec.get("group_size")) != ( original.in_features, original.out_features, original.bias is not None, GROUP_SIZE, ): raise ValueError(f"Linear metadata mismatch: {name}") qshape = list(packed_shape(original.out_features, original.in_features)) for suffix, shape, dtype in (("qweight", qshape, "U8"), ("scales", qshape[:2], "F32")): record = manifest["tensors"].get(f"{name}.{suffix}", {}) if record.get("shape") != shape or record.get("dtype") != dtype: raise ValueError(f"Quantized buffer metadata mismatch: {name}.{suffix}") replacements = {} for name, original in eligible.items(): if id(original) not in replacements: replacements[id(original)] = Q6Linear( original.in_features, original.out_features, original.bias is not None, device="meta", dtype=original.weight.dtype, ) parent_name, leaf = name.rsplit(".", 1) model.get_submodule(parent_name)._modules[leaf] = replacements[id(original)] expected = model.state_dict() actual = manifest["tensors"] if set(expected) != set(actual): raise ValueError(f"Checkpoint key mismatch: missing={sorted(set(expected)-set(actual))[:8]}, unexpected={sorted(set(actual)-set(expected))[:8]}") for name, value in expected.items(): if list(value.shape) != actual[name]["shape"]: raise ValueError(f"Checkpoint/model tensor shape mismatch: {name}") del expected # Capture alias groups before assignment; avoid overwriting tied tensors with # divergent checkpoint values. named_* with remove_duplicate=False is vital. aliases = {} for name, value in list(model.named_parameters(remove_duplicate=False)) + list(model.named_buffers(remove_duplicate=False)): aliases[name] = id(value) assigned = {} target_device = torch.device(device) loaded_count = 0 for filename in manifest["shards"]: with safe_open(str(path.parent / filename), framework="pt", device="cpu") as handle: for name in handle.keys(): tensor = handle.get_tensor(name) if name.endswith(".scales"): if not bool(torch.isfinite(tensor).all()) or not bool((tensor > 0).all()): raise ValueError(f"Nonfinite or nonpositive quantization scales: {name}") parent_name, _, leaf = name.rpartition(".") parent = model.get_submodule(parent_name) if parent_name else model alias_id = aliases[name] if alias_id in assigned: value = assigned[alias_id] if value.dtype != tensor.dtype or not torch.equal(value.detach().cpu(), tensor): raise ValueError(f"Conflicting values for aliased tensor: {name}") elif leaf in parent._parameters: value = nn.Parameter(tensor.to(target_device), requires_grad=False) assigned[alias_id] = value else: value = tensor.to(target_device) assigned[alias_id] = value if leaf in parent._parameters: parent._parameters[leaf] = value else: parent._buffers[leaf] = value loaded_count += 1 remaining_meta = [name for name, t in list(model.named_parameters()) + list(model.named_buffers()) if t.is_meta] if remaining_meta: raise ValueError(f"Registered tensors remain meta: {remaining_meta[:8]}") model.requires_grad_(False) model.eval() receipt = { "format": FORMAT, "profile": PROFILE, "manifest": str(path.resolve()), "manifest_sha256": sha256_file(path), "source": manifest["source"], "quantized_modules": len(definitions), "loaded_tensors": loaded_count, "shards": len(manifest["shards"]), "storage_bytes": sum(s["bytes"] for s in manifest["shards"].values()), "device": str(target_device), "elapsed_seconds": time.monotonic() - started, } return model, receipt