Prism-Q6 / loader /convert.py
atomtanstudio's picture
Add model card, Q6 manifest, loader package and conversion receipts
18e823d verified
Raw History Blame Contribute Delete
13.5 kB
#!/usr/bin/env python3
"""Stream an official Prism preview Safetensors into resumable custom-Q6 shards."""
from __future__ import annotations
import argparse
import hashlib
import json
import math
import os
from pathlib import Path
import struct
import time
import torch
from safetensors import safe_open
from safetensors.torch import save_file
from prism_quant.quant import FORMAT, VERSION, PROFILE, GROUP_SIZE, is_quantized_weight, packed_shape, quantize_weight
from prism_quant.quant_loader import DTYPES, sha256_file, verify_shards
def _log(event, **fields):
print(json.dumps({"event": event, **fields}, sort_keys=True), flush=True)
def _read_header(source: Path) -> dict:
with source.open("rb") as stream:
size_bytes = stream.read(8)
if len(size_bytes) != 8:
raise ValueError("Truncated Safetensors header")
size = struct.unpack("<Q", size_bytes)[0]
if not 2 <= size <= 16 * 1024 * 1024:
raise ValueError("Invalid or oversized Safetensors header")
header = json.loads(stream.read(size))
header.pop("__metadata__", None)
if not header:
raise ValueError("Empty source checkpoint")
for name, entry in header.items():
if entry.get("dtype") not in DTYPES or not isinstance(entry.get("shape"), list):
raise ValueError(f"Unsupported source tensor: {name}")
if any(type(n) is not int or n < 0 for n in entry["shape"]):
raise ValueError(f"Invalid source tensor shape: {name}")
return header
def _make_plan(header: dict, max_shard_bytes: int) -> tuple:
if max_shard_bytes <= 0:
raise ValueError("max_shard_bytes must be positive")
specs, modules, units = {}, {}, []
for name in sorted(header):
source_spec = header[name]
shape = source_spec["shape"]
if is_quantized_weight(name, shape):
if source_spec["dtype"] not in ("BF16", "F16", "F32"):
raise ValueError(f"Quantized Linear has unsupported source dtype: {name}")
module_name = name.removesuffix(".weight")
qshape = list(packed_shape(*shape))
names = [module_name + ".qweight", module_name + ".scales"]
if any(n in header for n in names):
raise ValueError(f"Source already contains quantized buffer names: {name}")
specs[names[0]] = {"shape": qshape, "dtype": "U8"}
specs[names[1]] = {"shape": qshape[:2], "dtype": "F32"}
modules[module_name] = {
"in_features": shape[1], "out_features": shape[0],
"bias": module_name + ".bias" in header, "group_size": GROUP_SIZE,
}
else:
names = [name]
specs[name] = {"shape": shape, "dtype": source_spec["dtype"]}
count = sum(math.prod(specs[n]["shape"]) * torch.tensor([], dtype=DTYPES[specs[n]["dtype"]]).element_size() for n in names)
if count > max_shard_bytes:
raise ValueError(f"Tensor storage unit {name} exceeds shard limit; increase --shard-mib")
units.append((name, names, count))
shards, current, current_size = [], [], 0
for unit in units:
if current and current_size + unit[2] > max_shard_bytes:
shards.append(current)
current, current_size = [], 0
current.append(unit)
current_size += unit[2]
if current:
shards.append(current)
for index, units in enumerate(shards, 1):
for _, names, _ in units:
for name in names:
specs[name]["shard"] = f"model-{index:05d}.safetensors"
return specs, modules, shards
def _verify_candidate(path: Path, specs: dict, identity: dict) -> dict:
with safe_open(str(path), framework="pt", device="cpu") as handle:
if handle.metadata() != identity or set(handle.keys()) != set(specs):
raise ValueError(f"Existing shard does not match this conversion: {path.name}")
for name, spec in specs.items():
view = handle.get_slice(name)
if view.get_shape() != spec["shape"] or view.get_dtype() != spec["dtype"]:
raise ValueError(f"Shard tensor mismatch: {name}")
return {"bytes": path.stat().st_size, "sha256": sha256_file(path)}
def _atomic_json(path: Path, value: dict):
partial = path.with_name(path.name + ".partial")
with partial.open("w") as stream:
json.dump(value, stream, sort_keys=True, indent=2)
stream.write("\n")
stream.flush()
os.fsync(stream.fileno())
os.replace(partial, path)
def convert_checkpoint(
source, destination, *, source_sha256: str, source_revision: str,
source_repo: str = "FrancisRing/Prism", max_shard_bytes: int = 512 * 1024 * 1024,
row_chunk: int = 256,
) -> Path:
"""CPU-only conversion. Completed shards can be reused after interruption.
The expected source hash is mandatory; even resumed conversion revalidates it.
Peak live tensor memory is bounded by one shard and one matrix's packed output,
plus the configured row scratch. No complete BF16 model is constructed.
"""
if row_chunk <= 0 or not source_revision or len(source_sha256) != 64:
raise ValueError("Provide a revision, a SHA-256, and a positive row_chunk")
try:
int(source_sha256, 16)
except ValueError as exc:
raise ValueError("Invalid source SHA-256") from exc
source_sha256 = source_sha256.lower()
source, destination = Path(source), Path(destination)
source_stat = source.stat()
_log("source_hash_start", path=str(source), bytes=source_stat.st_size)
started = time.monotonic()
if sha256_file(source) != source_sha256:
raise ValueError("Source SHA-256 does not match the pinned checkpoint")
header = _read_header(source)
specs, modules, shard_plan = _make_plan(header, max_shard_bytes)
source_info = {"repo": source_repo, "revision": source_revision, "sha256": source_sha256, "filename": source.name, "bytes": source_stat.st_size}
recipe = {"format": FORMAT, "version": VERSION, "profile": PROFILE, "group_size": GROUP_SIZE, "source": source_info, "tensors": specs, "quantized_modules": modules}
recipe_hash = hashlib.sha256(json.dumps(recipe, sort_keys=True, separators=(",", ":")).encode()).hexdigest()
destination.mkdir(parents=True, exist_ok=True)
final_path = destination / "manifest.json"
if final_path.exists():
existing = verify_shards(final_path)
if existing.get("recipe_sha256") != recipe_hash:
raise ValueError("Destination contains a different completed conversion")
_log("conversion_already_complete", manifest=str(final_path))
return final_path
shard_records = {}
sample_errors = {}
completed_tensors = 0
# safe_open validates offsets and byte lengths independently of our plan.
with safe_open(str(source), framework="pt", device="cpu") as handle:
if set(handle.keys()) != set(header):
raise ValueError("Source header keys changed")
for index, units in enumerate(shard_plan, 1):
filename = f"model-{index:05d}.safetensors"
path = destination / filename
shard_specs = {name: specs[name] for _, names, _ in units for name in names}
identity = {"format": FORMAT, "version": str(VERSION), "recipe_sha256": recipe_hash, "source_sha256": source_sha256, "shard": str(index)}
receipt_path = destination / (filename + ".receipt.json")
if path.exists():
record = _verify_candidate(path, shard_specs, identity)
if not receipt_path.exists():
# A crash after atomic shard rename but before its receipt is
# recoverable by recomputing this bounded shard from source.
_log("shard_rebuild_without_receipt", shard=filename)
else:
old = json.loads(receipt_path.read_text())
if old.get("record") != record or old.get("recipe_sha256") != recipe_hash:
raise ValueError(f"Existing shard receipt/hash mismatch: {filename}")
shard_records[filename] = record
sample_errors.update(old.get("sample_errors", {}))
completed_tensors += len(units)
_log("shard_resumed", shard=filename, completed_source_tensors=completed_tensors, total_source_tensors=len(header))
continue
tensors = {}
shard_errors = {}
for name, output_names, _ in units:
shape = header[name]["shape"]
if is_quantized_weight(name, shape):
view = handle.get_slice(name)
qshape = packed_shape(*shape)
packed = torch.empty(qshape, dtype=torch.uint8, device="cpu")
scales = torch.empty(qshape[:2], dtype=torch.float32, device="cpu")
for row in range(0, shape[0], row_chunk):
stop = min(row + row_chunk, shape[0])
weights = view[row:stop, :]
q, s = quantize_weight(weights)
packed[row:stop].copy_(q)
scales[row:stop].copy_(s)
# An actual-weight diagnostic on a bounded sample, not a
# claim about end-to-end perceptual quality.
if row == 0:
from prism_quant.quant import dequantize_weight
sample_rows = min(stop, 8)
original = weights[:sample_rows].float()
decoded = dequantize_weight(q[:sample_rows], s[:sample_rows], (sample_rows, shape[1]), dtype=torch.float32)
mse = (original - decoded).square().mean().item()
power = original.square().mean().item()
shard_errors[name] = {"sample_rows": sample_rows, "rmse": math.sqrt(mse), "relative_rmse": math.sqrt(mse / power) if power else 0.0}
del weights, q, s
tensors[output_names[0]], tensors[output_names[1]] = packed, scales
else:
tensor = handle.get_tensor(name)
if tensor.is_floating_point() and not bool(torch.isfinite(tensor).all()):
raise ValueError(f"Nonfinite retained tensor: {name}")
tensors[name] = tensor.contiguous()
completed_tensors += 1
_log("tensor_complete", tensor=name, completed_source_tensors=completed_tensors, total_source_tensors=len(header), elapsed_seconds=round(time.monotonic() - started, 3))
partial = path.with_name(path.name + ".partial")
save_file(tensors, str(partial), metadata=identity)
del tensors
record = _verify_candidate(partial, shard_specs, identity)
with partial.open("rb") as stream:
os.fsync(stream.fileno())
os.replace(partial, path)
_atomic_json(receipt_path, {"recipe_sha256": recipe_hash, "record": record, "sample_errors": shard_errors})
shard_records[filename] = record
sample_errors.update(shard_errors)
_log("shard_complete", shard=filename, shards_total=len(shard_plan), bytes=record["bytes"], sha256=record["sha256"])
final_stat = source.stat()
if (source_stat.st_size, source_stat.st_mtime_ns, source_stat.st_ino) != (final_stat.st_size, final_stat.st_mtime_ns, final_stat.st_ino):
raise ValueError("Source file changed during conversion")
manifest = {**recipe, "recipe_sha256": recipe_hash, "shards": shard_records, "sample_errors": sample_errors,
"conversion": {"row_chunk": row_chunk, "max_shard_bytes": max_shard_bytes, "elapsed_seconds": time.monotonic() - started}}
# Publish only after all shards are complete and verified, including resumed
# shards. The temporary manifest is itself valid for the strict validator.
pending = destination / "manifest.pending.json"
_atomic_json(pending, manifest)
verify_shards(pending)
os.replace(pending, final_path)
_log("conversion_complete", manifest=str(final_path), quantized_modules=len(modules), storage_bytes=sum(v["bytes"] for v in shard_records.values()))
return final_path
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--source", required=True, type=Path)
parser.add_argument("--output-dir", required=True, type=Path)
parser.add_argument("--source-sha256", required=True)
parser.add_argument("--source-revision", required=True)
parser.add_argument("--source-repo", default="FrancisRing/Prism")
parser.add_argument("--shard-mib", type=int, default=512)
parser.add_argument("--row-chunk", type=int, default=256)
parser.add_argument("--threads", type=int, default=8)
args = parser.parse_args()
if args.threads <= 0:
parser.error("--threads must be positive")
torch.set_num_threads(args.threads)
convert_checkpoint(args.source, args.output_dir, source_sha256=args.source_sha256,
source_revision=args.source_revision, source_repo=args.source_repo,
max_shard_bytes=args.shard_mib * 1024 * 1024, row_chunk=args.row_chunk)
if __name__ == "__main__":
main()