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