Download code/scripts/check_w4.py from EndlessChasing/Mamb2_8B_W4A16_Recall: direct link, hf CLI and curl.
- Browser
- Download file 7.09 kB
-
https://huggingface.co/EndlessChasing/Mamb2_8B_W4A16_Recall/resolve/main/code/scripts/check_w4.py
- Command line
-
hf download hf://EndlessChasing/Mamb2_8B_W4A16_Recall/code/scripts/check_w4.py
-
curl -L -o check_w4.py https://huggingface.co/EndlessChasing/Mamb2_8B_W4A16_Recall/resolve/main/code/scripts/check_w4.py
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() | |