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