import struct import sys from pathlib import Path ROOT = Path(__file__).resolve().parents[1] if str(ROOT) not in sys.path: sys.path.insert(0, str(ROOT)) import torch from b1_tensor_core.ops import b1_pack from scripts.quantize_outer_int4 import unpack_int4_signed from src.r4t.b1_diffusion import B1EDMDenoiser CKPT_PATH = ROOT / "checkpoints" / "champion_b1_consistency_1step_qat.pt" OUTPUT_BIN_PATH = ROOT / "models" / "fanout_1bit_weights.bin" def get_packed(weight: torch.Tensor) -> torch.Tensor: """Takes [out_features, in_features] and returns packed [out_features, in_features // 32] uint32.""" bipolar = torch.where(weight >= 0, 1.0, -1.0).float() return b1_pack(bipolar) def main(): print(f"Loading checkpoint from: {CKPT_PATH}") ckpt = torch.load(CKPT_PATH, map_location="cpu", weights_only=False) config = ckpt["config"] model = B1EDMDenoiser(config, backend="tc", pure_1bit=False) state = model.state_dict() if "weights" in ckpt: for k, v in ckpt["weights"].items(): if k in state: state[k].copy_(v) if "int4_outer" in ckpt: for k, d in ckpt["int4_outer"].items(): state[k].copy_(unpack_int4_signed(d["packed"], d["scale"])) elif "model_state_dict" in ckpt: model.load_state_dict(ckpt["model_state_dict"], strict=False) OUTPUT_BIN_PATH.parent.mkdir(parents=True, exist_ok=True) f = open(OUTPUT_BIN_PATH, "wb") # Magic & Header f.write(b"B1FO") # header: [version, hidden_dim, embedding_dim, num_layers, num_heads, mlp_dim, seq_len] header = struct.pack( "<7I", 1, config.hidden_dim, config.embedding_dim, config.layers, config.heads, config.mlp_dim, config.sequence_length, ) f.write(header) def write_tensor_f32(t: torch.Tensor): arr = t.detach().cpu().float().contiguous().numpy() f.write(arr.tobytes()) def write_tensor_u32(t: torch.Tensor): arr = t.detach().cpu().contiguous().numpy() f.write(arr.tobytes()) def get_layer_packed(prefix: str, in_feat: int) -> torch.Tensor: if f"{prefix}.packed_weight" in ckpt["weights"]: return ckpt["weights"][f"{prefix}.packed_weight"] elif f"{prefix}.weight" in state: return get_packed(state[f"{prefix}.weight"]) else: raise KeyError(f"Weight not found for {prefix}") def get_layer_bias(prefix: str, out_feat: int) -> torch.Tensor: k = f"{prefix}.bias" if k in ckpt.get("weights", {}): return ckpt["weights"][k] elif k in state and state[k] is not None: return state[k] return torch.zeros(out_feat, dtype=torch.float32) # 1. Positional embedding pos = state["backbone.position"].squeeze(0) # [10, 512] write_tensor_f32(pos) # 2. Time MLP write_tensor_f32(state["backbone.time_mlp.0.weight"]) # [512, 512] write_tensor_f32(state["backbone.time_mlp.0.bias"]) # [512] write_tensor_f32(state["backbone.time_mlp.2.weight"]) # [512, 512] write_tensor_f32(state["backbone.time_mlp.2.bias"]) # [512] # 3. Outer projections write_tensor_f32(state["backbone.input_projection.weight"]) # [512, 768] write_tensor_f32(state["backbone.input_projection.bias"]) # [512] write_tensor_f32(state["backbone.query_projection.weight"]) # [512, 768] write_tensor_f32(state["backbone.query_projection.bias"]) # [512] write_tensor_f32(state["backbone.output_projection.weight"]) # [768, 512] write_tensor_f32(state["backbone.output_projection.bias"]) # [768] # 4. Final Norm write_tensor_f32(state["backbone.final_norm.weight"]) write_tensor_f32(state["backbone.final_norm.bias"]) # 5. Transformer Layers for i in range(config.layers): p = f"backbone.layers.{i}" # Norms write_tensor_f32(state[f"{p}.norm1.weight"]) write_tensor_f32(state[f"{p}.norm1.bias"]) write_tensor_f32(state[f"{p}.norm2.weight"]) write_tensor_f32(state[f"{p}.norm2.bias"]) write_tensor_f32(state[f"{p}.norm3.weight"]) write_tensor_f32(state[f"{p}.norm3.bias"]) # Self-Attention write_tensor_u32(get_layer_packed(f"{p}.self_attn.q_proj.linear", config.hidden_dim)) write_tensor_f32(get_layer_bias(f"{p}.self_attn.q_proj.linear", config.hidden_dim)) write_tensor_u32(get_layer_packed(f"{p}.self_attn.k_proj.linear", config.hidden_dim)) write_tensor_f32(get_layer_bias(f"{p}.self_attn.k_proj.linear", config.hidden_dim)) write_tensor_u32(get_layer_packed(f"{p}.self_attn.v_proj.linear", config.hidden_dim)) write_tensor_f32(get_layer_bias(f"{p}.self_attn.v_proj.linear", config.hidden_dim)) write_tensor_u32(get_layer_packed(f"{p}.self_attn.out_proj.linear", config.hidden_dim)) write_tensor_f32(get_layer_bias(f"{p}.self_attn.out_proj.linear", config.hidden_dim)) # Cross-Attention write_tensor_u32(get_layer_packed(f"{p}.cross_attn.q_proj.linear", config.hidden_dim)) write_tensor_f32(get_layer_bias(f"{p}.cross_attn.q_proj.linear", config.hidden_dim)) write_tensor_u32(get_layer_packed(f"{p}.cross_attn.k_proj.linear", config.hidden_dim)) write_tensor_f32(get_layer_bias(f"{p}.cross_attn.k_proj.linear", config.hidden_dim)) write_tensor_u32(get_layer_packed(f"{p}.cross_attn.v_proj.linear", config.hidden_dim)) write_tensor_f32(get_layer_bias(f"{p}.cross_attn.v_proj.linear", config.hidden_dim)) write_tensor_u32(get_layer_packed(f"{p}.cross_attn.out_proj.linear", config.hidden_dim)) write_tensor_f32(get_layer_bias(f"{p}.cross_attn.out_proj.linear", config.hidden_dim)) # MLP write_tensor_u32(get_layer_packed(f"{p}.mlp.fc1.linear", config.hidden_dim)) write_tensor_f32(get_layer_bias(f"{p}.mlp.fc1.linear", config.mlp_dim)) write_tensor_u32(get_layer_packed(f"{p}.mlp.fc2.linear", config.mlp_dim)) write_tensor_f32(get_layer_bias(f"{p}.mlp.fc2.linear", config.hidden_dim)) f.close() size_mb = OUTPUT_BIN_PATH.stat().st_size / (1024 * 1024) print(f"Exported C++ Flat Binary Model Weights: {OUTPUT_BIN_PATH} ({size_mb:.2f} MB)") if __name__ == "__main__": main()