File size: 9,169 Bytes
18e823d | 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 | """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
|