EndlessChasing's picture
Release complete independent W4A16 Mamba2-8B with Resurface adapter
cf0f656 verified
Raw History Blame Contribute Delete
7.09 kB
#!/usr/bin/env python3
"""Check the codec or every file and decoded hash in a completed W4 package."""
from __future__ import annotations
import argparse
import copy
import hashlib
import json
import sys
import tempfile
from pathlib import Path
import torch
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from mamba2_recall import runtime, w4
def _raises(call, exceptions=(ValueError, TypeError)):
try:
call()
except exceptions:
return
raise AssertionError("Expected invalid input to be rejected")
def self_test():
"""Exercise independent bit patterns, edges, cross-row chunks and corruption."""
torch.manual_seed(20260928)
codes = torch.arange(16, dtype=torch.uint8).reshape(2, 8)
packed = w4.pack_nibbles(codes)
expected = torch.tensor([[0x10, 0x32, 0x54, 0x76], [0x98, 0xba, 0xdc, 0xfe]], dtype=torch.uint8)
assert torch.equal(packed, expected)
assert torch.equal(w4.unpack_nibbles(expected, 8), codes)
odd = torch.tensor([[0, 15, 3], [9, 2, 14]], dtype=torch.uint8)
assert torch.equal(w4.unpack_nibbles(w4.pack_nibbles(odd), 3), odd)
_raises(lambda: w4.pack_nibbles(torch.tensor([[16]], dtype=torch.uint8)))
_raises(lambda: w4.unpack_nibbles(torch.tensor([[0xff]], dtype=torch.uint8), 1))
for fill in (0., 1.25, -7.5, torch.finfo(torch.float16).smallest_normal / 1024):
weights = torch.full((3, 257), fill, dtype=torch.float16)
scales, offsets, codes, stats = w4.quantize_groups(weights)
assert torch.equal(w4.decode_groups(scales, offsets, codes), weights)
assert stats["squared_error"] == 0
assert bool((scales == 0).all())
for value in (float("nan"), float("inf"), -float("inf"), 100000.):
_raises(lambda: w4.quantize_groups(torch.tensor([[value]], dtype=torch.float32)))
weights = (torch.randn(7, 259) * .07).half()
weights[0, :4] = torch.tensor([0., 0.000000059604645, -0.000000059604645, 0.])
scales, offsets, codes, stats = w4.quantize_groups(weights)
decoded = w4.decode_groups(scales, offsets, codes)
actual_sse = (decoded.float() - weights.float()).square().double().sum().item()
assert abs(actual_sse - stats["squared_error"]) <= 1e-6 * max(1., actual_sse)
assert stats["squared_error"] <= stats["unclipped_squared_error"] + 1e-8
assert sum(stats["candidate_group_counts"]) == stats["group_count"]
# Saturated FP16 endpoints must yield a finite selected representation.
extreme = torch.tensor([[-65504., 65504.] * 64], dtype=torch.float16)
sx, ox, cx, _ = w4.quantize_groups(extreme)
assert bool(torch.isfinite(w4.decode_groups(sx, ox, cx)).all())
with tempfile.TemporaryDirectory(prefix="mamba2-w4-check-") as temporary:
directory = Path(temporary)
path = directory / "0000.w4bin"
entry = w4.write_tensor(path, weights, "w4_affine_f16", chunk_rows=2, reserve_bytes=0)
restored = w4.read_tensor(path, chunk_rows=3)
assert torch.equal(restored, decoded)
assert entry["decoded_sha256"] == hashlib.sha256(w4._bytes(decoded)).hexdigest()
w4.verify_tensor_entry(directory, entry)
fp16_path = directory / "0001.w4bin"
fp16_source = torch.randn(3, 4, 5).bfloat16()
fp16_entry = w4.write_tensor(fp16_path, fp16_source, "fp16", reserve_bytes=0)
assert torch.equal(w4.read_tensor(fp16_path, chunk_rows=2), fp16_source.half())
w4.verify_tensor_entry(directory, fp16_entry)
data = path.read_bytes()
for length in (0, w4.PREFIX.size - 1, w4.PREFIX.size + 1, len(data) - 1):
bad = directory / "truncated.w4bin"
bad.write_bytes(data[:length])
_raises(lambda: w4.read_tensor(bad))
bad.write_bytes(data + b"\x00")
_raises(lambda: w4.read_tensor(bad))
corrupt = bytearray(data)
corrupt[-2] ^= 1
path.write_bytes(corrupt)
_raises(lambda: w4.verify_tensor_entry(directory, entry))
path.write_bytes(data)
# Manifest checks fail before any model allocation for wrong completion,
# source, config, algorithm, counts or coverage.
base = {"format": w4.FORMAT, "version": 1, "complete": True,
"source_checkpoint_sha256": runtime.SOURCE_CHECKPOINT_SHA256,
"model_config": copy.deepcopy(runtime.MODEL_CONFIG),
"quantization": copy.deepcopy(w4.QUANTIZATION),
"parameter_count": w4.PARAMETER_COUNT, "tensor_count": w4.TENSOR_COUNT,
"w4_tensor_count": w4.W4_TENSOR_COUNT, "fp16_tensor_count": 393,
"tensors": {}, "tensor_bytes": 0}
for key, value in (("complete", False), ("source_checkpoint_sha256", "0" * 64),
("model_config", {}), ("quantization", {}), ("tensor_count", 506),
("parameter_count", w4.PARAMETER_COUNT - 1)):
invalid = copy.deepcopy(base)
invalid[key] = value
_raises(lambda: w4.validate_manifest(invalid, expected_shapes={}))
_raises(lambda: w4.validate_manifest(base, expected_shapes={}))
return {"passed": True, "checks": ["all_nibble_codes_and_byte_order", "odd_columns_and_padding",
"zero_constant_subnormal_groups", "nonfinite_rejection", "fp16_endpoint_overflow",
"serialized_mse_selection", "cross_chunk_row_roundtrip", "fp16_roundtrip",
"truncation_and_trailing_bytes", "file_and_decoded_hashes", "strict_invalid_manifest"]}
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--self-test", action="store_true")
parser.add_argument("--package")
parser.add_argument("--cpu-threads", type=int, default=8)
parser.add_argument("--load-model", action="store_true", help="Also instantiate and strictly load the FP16 native model")
parser.add_argument("--device", default="cuda")
args = parser.parse_args()
if not args.self_test and not args.package:
parser.error("Choose --self-test and/or --package")
if args.cpu_threads < 1:
parser.error("--cpu-threads must be positive")
torch.set_num_threads(args.cpu_threads)
if args.self_test:
print(json.dumps({"event": "codec_self_test", **self_test()}), flush=True)
if args.package:
manifest = w4.read_manifest(args.package)
for index, (name, entry) in enumerate(manifest["tensors"].items(), 1):
w4.verify_tensor_entry(args.package, entry)
print(json.dumps({"event": "tensor_verified", "index": index, "tensor": name}), flush=True)
result = {"event": "package_verified", "tensor_count": len(manifest["tensors"]),
"manifest_sha256": runtime.sha256_file(Path(args.package) / "manifest.json"),
"tensor_bytes": manifest["tensor_bytes"]}
if args.load_model:
model = w4.load_w4_model(args.package, device=args.device)
result["model_receipt"] = model._package_receipt
print(json.dumps(result), flush=True)
if __name__ == "__main__":
main()