#!/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(" 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()