Download quantize.py from M1n1A1/MiniAI-TITS2: direct link, hf CLI and curl.
- Browser
- Download file 3.17 kB
-
https://huggingface.co/M1n1A1/MiniAI-TITS2/resolve/main/quantize.py
- Command line
-
hf download hf://M1n1A1/MiniAI-TITS2/quantize.py
-
curl -L -o quantize.py https://huggingface.co/M1n1A1/MiniAI-TITS2/resolve/main/quantize.py
3.17 kB
| """ | |
| T.I.T.S.2 checkpoint tools: export to safetensors, and quantize to 1 bit for laughs. | |
| safetensors : the sane, safe, memory-mappable format everyone uses now. | |
| 1-bit : every weight becomes sign(w), with one per-row scale (BitNet-style | |
| absmean). 16x smaller than fp16. Real 1-bit models are *trained* | |
| that way; crushing a finished fp32 model to 1 bit post-hoc is not | |
| expected to survive, which is the entertainment. | |
| Usage: | |
| python quantize.py --checkpoint checkpoints2/tits2_epoch24.pt --out_dir exports | |
| """ | |
| import argparse | |
| import os | |
| import torch | |
| from safetensors.torch import save_file | |
| SKIP = ("norm", "pos", "bias", "x_embed", "proj_out", "ada") # tiny or structural: leave alone | |
| def one_bit(w): | |
| """w -> (packed sign bits, per-row scale). Reconstruct: sign * scale.""" | |
| scale = w.abs().mean(dim=1, keepdim=True) # absmean per output row | |
| signs = (w >= 0) | |
| packed = torch.zeros((w.shape[0], (w.shape[1] + 7) // 8), dtype=torch.uint8) | |
| flat = signs.to(torch.uint8) | |
| for bit in range(8): | |
| chunk = flat[:, bit::8] | |
| packed[:, : chunk.shape[1]] |= chunk << bit | |
| return packed, scale.to(torch.float16) | |
| def unpack(packed, scale, out_features): | |
| bits = torch.zeros((packed.shape[0], packed.shape[1] * 8), dtype=torch.uint8) | |
| for bit in range(8): | |
| bits[:, bit::8] = (packed >> bit) & 1 | |
| signs = bits[:, :out_features].float() * 2 - 1 # {0,1} -> {-1,+1} | |
| return signs * scale.float() | |
| def main(): | |
| p = argparse.ArgumentParser() | |
| p.add_argument("--checkpoint", default="checkpoints2/tits2_epoch24.pt") | |
| p.add_argument("--out_dir", default="exports") | |
| args = p.parse_args() | |
| os.makedirs(args.out_dir, exist_ok=True) | |
| ck = torch.load(args.checkpoint, map_location="cpu", weights_only=False) | |
| sd = ck["ema_state_dict"] | |
| meta = {k: str(ck[k]) for k in ("dim", "depth", "heads", "latent_size", "image_size", "step")} | |
| fp16 = {k: v.half().contiguous() for k, v in sd.items()} | |
| fp16_path = os.path.join(args.out_dir, "tits2_ema_fp16.safetensors") | |
| save_file(fp16, fp16_path, metadata=meta) | |
| out, n_quant, bits_before, bits_after = {}, 0, 0, 0 | |
| for k, v in sd.items(): | |
| if v.dim() == 2 and not any(s in k for s in SKIP): | |
| packed, scale = one_bit(v.float()) | |
| out[k + ".packed"] = packed | |
| out[k + ".scale"] = scale | |
| out[k + ".shape"] = torch.tensor(list(v.shape), dtype=torch.int32) | |
| n_quant += 1 | |
| bits_before += v.numel() * 16 | |
| bits_after += v.numel() + scale.numel() * 16 | |
| else: | |
| out[k] = v.half().contiguous() | |
| onebit_path = os.path.join(args.out_dir, "tits2_1bit.safetensors") | |
| save_file(out, onebit_path, metadata={**meta, "quant": "1bit-absmean-per-row"}) | |
| mb = lambda p: os.path.getsize(p) / 1e6 | |
| print(f"fp16 safetensors : {mb(fp16_path):7.1f} MB {fp16_path}") | |
| print(f"1-bit safetensors: {mb(onebit_path):7.1f} MB {onebit_path} ({n_quant} matrices quantized)") | |
| print(f"quantized weights: {bits_before/8e6:.1f} MB -> {bits_after/8e6:.1f} MB") | |
| if __name__ == "__main__": | |
| main() | |