Prism-Q6 / loader /prism_quant /quant_loader.py
atomtanstudio's picture
Add model card, Q6 manifest, loader package and conversion receipts
18e823d verified
Raw History Blame Contribute Delete
9.17 kB
"""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