""" 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()