MiniAI-TITS2 / quantize.py
Codeminute's picture
T.I.T.S.2 — 93M DiT, rectified flow, 256px, 24 epochs on cleaned CC12M
2ca14fe verified
Raw History Blame Contribute Delete
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()