File size: 6,334 Bytes
f08972a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 | 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()
|