Download loader/prism_quant/quant_loader.py from atomtanstudio/Prism-Q6: direct link, hf CLI and curl.
- Browser
- Download file 9.17 kB
-
https://huggingface.co/atomtanstudio/Prism-Q6/resolve/main/loader/prism_quant/quant_loader.py
- Command line
-
hf download hf://atomtanstudio/Prism-Q6/loader/prism_quant/quant_loader.py
-
curl -L -o quant_loader.py https://huggingface.co/atomtanstudio/Prism-Q6/resolve/main/loader/prism_quant/quant_loader.py
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 | |