fanout-diffusion / scripts /export_cpp_weights.py
dejanseo's picture
Add training pipelines, consistency distillation scripts, and interactive dashboard server
f08972a verified
Raw History Blame Contribute Delete
6.33 kB
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()