File size: 11,456 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
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
import math
import sys
import time
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
import torch.nn.functional as F
from sentence_transformers import SentenceTransformer
from torch.utils.data import DataLoader, TensorDataset

from scripts.fast_b1_inference import FastB1Denoiser, fast_sample_edm_8step
from src.r4t.b1_diffusion import B1EDMDenoiser
from src.r4t.config import DiffusionConfig
from src.r4t.diffusion import diffusion_loss, sample_edm
from src.r4t.journal import ExperimentJournal

CHAMPION_CKPT = Path("checkpoints/b1_tc_10ep_champion.pt")
INT4_EXPORT_PATH = Path("checkpoints/champion_b1_tc_int4_outer.pt")
DATA_PATH = Path("data/diffusion_dataset_540k.pt")
TAXONOMY_PATH = Path("data/taxonomy_embeddings.pt")


def pack_int4_signed(tensor: torch.Tensor):
    """
    Symmetric per-channel INT4 quantization:
    Quantizes float tensor in range [-8, 7] and packs pairs of 4-bit nibbles into uint8.
    """
    # Per-row scale: [M, 1]
    max_val = tensor.abs().max(dim=-1, keepdim=True).values.clamp_min(1e-8)
    scale = max_val / 7.0  # Range -7 to +7 (or -8 to 7)

    q = torch.clamp(torch.round(tensor / scale), -8, 7).to(torch.int8)

    # Convert signed 4-bit to unsigned 4-bit [0..15]
    q_u = (q & 0x0F).to(torch.uint8)

    # Pack adjacent elements (dim=-1 must be even)
    # low nibble = even, high nibble = odd
    q_even = q_u[..., 0::2]
    q_odd = q_u[..., 1::2]
    packed = (q_odd << 4) | q_even

    return packed, scale.to(torch.float16)


def unpack_int4_signed(packed: torch.Tensor, scale: torch.Tensor):
    """Unpacks pairs of 4-bit nibbles from uint8 back to float16 tensor."""
    q_even = (packed & 0x0F).to(torch.int8)
    q_odd = ((packed >> 4) & 0x0F).to(torch.int8)

    # Sign extend from 4-bit to 8-bit
    q_even = torch.where(q_even >= 8, q_even - 16, q_even)
    q_odd = torch.where(q_odd >= 8, q_odd - 16, q_odd)

    M = packed.shape[0]
    K = packed.shape[1] * 2
    unpacked = torch.empty((M, K), dtype=torch.float16, device=packed.device)
    unpacked[:, 0::2] = q_even.to(torch.float16)
    unpacked[:, 1::2] = q_odd.to(torch.float16)

    return unpacked * scale.to(unpacked.device)


def evaluate_qualitative(model, embedder, tax_emb, tax_names, device):
    test_queries = [
        "quantum computing algorithms for cryptography",
        "renewable energy storage systems and solar cells",
        "deep neural networks for medical image diagnostics",
    ]
    total_alignment = 0.0
    total_diversity = 0.0

    model.eval()
    with torch.no_grad():
        for q_text in test_queries:
            q_emb = embedder.encode([q_text], convert_to_tensor=True, device=device).float()
            q_emb = F.normalize(q_emb, dim=-1)

            subq_traj = sample_edm(model, q_emb, sampling_steps=8, cfg_strength=0.1)
            subq_emb = F.normalize(subq_traj[0], dim=-1)

            # Alignment
            sim_prompt = (subq_emb @ q_emb.squeeze(0)).mean().item()
            total_alignment += sim_prompt

            # Diversity
            sim_matrix = subq_emb @ subq_emb.T
            L = sim_matrix.shape[0]
            mask = ~torch.eye(L, dtype=torch.bool, device=device)
            pairwise_div = 1.0 - sim_matrix[mask].mean().item()
            total_diversity += pairwise_div

    avg_align = total_alignment / len(test_queries)
    avg_div = total_diversity / len(test_queries)
    return avg_align, avg_div


def main():
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    print(f"Device: {device} ({torch.cuda.get_device_name(0)})")

    print(f"Loading champion checkpoint from {CHAMPION_CKPT}...")
    ckpt = torch.load(CHAMPION_CKPT, map_location=device, weights_only=False)
    config = ckpt["config"]

    model = B1EDMDenoiser(config, backend="tc", pure_1bit=False).to(device)
    if "ema_state_dict" in ckpt and "shadow" in ckpt["ema_state_dict"]:
        shadow = ckpt["ema_state_dict"]["shadow"]
        model.load_state_dict({k: shadow[k].to(device) for k in shadow})
    else:
        model.load_state_dict(ckpt["model_state_dict"])

    model.eval()
    model.freeze_for_inference()

    # Outer adapters to quantize to INT4
    outer_keys = [
        "backbone.input_projection.weight",
        "backbone.output_projection.weight",
        "backbone.query_projection.weight",
        "backbone.time_mlp.0.weight",
        "backbone.time_mlp.2.weight",
    ]

    export_dict = {
        "config": config,
        "weights": {},
        "int4_outer": {},
    }

    state = model.state_dict()
    total_unpacked_bytes = 0
    total_int4_bytes = 0

    print("\nQuantizing Outer Adapters to Symmetric INT4:")
    print("--------------------------------------------------------------------------------")
    print(f"{'Layer':<40} | {'Orig FP16':<12} | {'INT4 Packed':<12} | {'MSE Error':<10}")
    print("--------------------------------------------------------------------------------")

    for k, v in state.items():
        if k in outer_keys:
            orig_bytes = v.numel() * 2
            total_unpacked_bytes += orig_bytes

            packed, scale = pack_int4_signed(v.float())
            recon = unpack_int4_signed(packed, scale)
            mse = F.mse_loss(recon.float(), v.float()).item()

            export_dict["int4_outer"][k] = {
                "packed": packed.cpu(),
                "scale": scale.cpu(),
            }
            int4_bytes = packed.numel() + scale.numel() * 2
            total_int4_bytes += int4_bytes

            print(f"{k:<40} | {orig_bytes/1024:>8.1f} KB | {int4_bytes/1024:>8.1f} KB | {mse:.2e}")
        elif "packed_weight" in k:
            export_dict["weights"][k] = v.cpu()
            int4_bytes = v.numel() * v.element_size()
            total_int4_bytes += int4_bytes
            total_unpacked_bytes += int4_bytes
        elif "weight" in k and any(proj in k for proj in ["self_attn", "cross_attn", "mlp"]):
            # Skip uncompressed latent FP32 weights of B1Linear layers
            continue
        else:
            # Biases, LayerNorms, positional embeddings (FP16)
            v_fp16 = v.to(torch.float16).cpu()
            export_dict["weights"][k] = v_fp16
            int4_bytes = v_fp16.numel() * 2
            total_int4_bytes += int4_bytes
            total_unpacked_bytes += int4_bytes

    print("--------------------------------------------------------------------------------")
    print(f"Original Hybrid Model Size:  {total_unpacked_bytes / (1024*1024):.2f} MB")
    print(f"INT4 Outer Quantized Model:  {total_int4_bytes / (1024*1024):.2f} MB (raw tensors)")

    # Save to disk
    torch.save(export_dict, INT4_EXPORT_PATH)
    file_size_bytes = INT4_EXPORT_PATH.stat().st_size
    print(f"\nSaved INT4 Deployment Checkpoint: {INT4_EXPORT_PATH}")
    print(f"File Size on Disk: {file_size_bytes / 1024:.1f} KB ({file_size_bytes / (1024*1024):.2f} MB)")

    # Apply reconstructed weights back to model to test validation loss & fidelity
    for k in outer_keys:
        p = export_dict["int4_outer"][k]["packed"].to(device)
        s = export_dict["int4_outer"][k]["scale"].to(device)
        recon = unpack_int4_signed(p, s)
        state[k].copy_(recon)

    print("\nValidating Fidelity of INT4-Quantized Model on Dataset...")
    dataset_dict = torch.load(DATA_PATH, map_location="cpu", weights_only=False)
    queries = dataset_dict["query_embeddings"].float()
    targets = dataset_dict["targets"].float()
    N, L, D = targets.shape
    n_train = int(0.9 * N)
    val_queries, val_targets = queries[n_train:], targets[n_train:]
    val_loader = DataLoader(TensorDataset(val_queries, val_targets), batch_size=128, shuffle=False)

    val_loss_total = 0.0
    val_gen = torch.Generator(device=device).manual_seed(1337)
    with torch.no_grad():
        for b_queries, b_targets in val_loader:
            b_queries = b_queries.to(device, non_blocking=True)
            b_targets = b_targets.to(device, non_blocking=True)

            sims = torch.einsum("bd,bld->bl", F.normalize(b_queries, dim=-1), F.normalize(b_targets, dim=-1))
            sorted_idx = torch.argsort(sims, dim=1, descending=True)
            b_targets = torch.gather(b_targets, 1, sorted_idx.unsqueeze(-1).expand(-1, -1, D))

            v_loss = diffusion_loss(model, b_targets, b_queries, generator=val_gen)
            val_loss_total += v_loss.item() * len(b_queries)

    val_loss = val_loss_total / len(val_queries)
    print(f"Validation Loss after INT4 Outer Quantization: {val_loss:.4f} (Baseline FP16: 0.6885)")

    print("\nEvaluating Qualitative Decoding (EmbeddingGemma)...")
    embedder = SentenceTransformer("google/embeddinggemma-300m", model_kwargs={"torch_dtype": torch.bfloat16}, device=device)
    tax_dict = torch.load(TAXONOMY_PATH, map_location=device, weights_only=False)
    tax_emb = tax_dict["embeddings"].to(device).float()
    tax_names = tax_dict["names"]

    align, div = evaluate_qualitative(model, embedder, tax_emb, tax_names, device)
    print(f"INT4 Outer Model: Prompt Alignment = {align:.3f} (FP16: 0.324) | Diversity = {div:.3f} (FP16: 0.844)")

    # Benchmark latency
    fast_model = FastB1Denoiser(model)
    dummy_q = torch.randn(1, 768, device=device)
    dummy_q = F.normalize(dummy_q, dim=-1)

    # CUDA Graph
    g_stream = torch.cuda.Stream()
    g_stream.wait_stream(torch.cuda.current_stream())
    with torch.cuda.stream(g_stream):
        for _ in range(3):
            _ = fast_sample_edm_8step(fast_model, dummy_q)
    torch.cuda.current_stream().wait_stream(g_stream)

    g = torch.cuda.CUDAGraph()
    with torch.cuda.graph(g, stream=g_stream):
        _ = fast_sample_edm_8step(fast_model, dummy_q)

    torch.cuda.synchronize()
    times = []
    for _ in range(100):
        t0 = time.perf_counter()
        g.replay()
        torch.cuda.synchronize()
        times.append((time.perf_counter() - t0) * 1000.0)

    avg_ms = sum(times) / len(times)
    p95_ms = sorted(times)[int(len(times) * 0.95)]
    qps = 1000.0 / avg_ms

    print(f"\nCUDA Graph 8-Step Heun Latency: {avg_ms:.2f} ms (P95: {p95_ms:.2f} ms, {qps:.1f} QPS)")

    # Log to journal
    journal = ExperimentJournal()
    tracker = journal.start_run(
        name="Champion B1-TC: INT4 Outer Quantization (1.5 MB)",
        experiment_name="1-Bit Tensor Core Innovation",
        task_type="diffusion",
        config={
            "core": "1-bit_ptx_mma",
            "outer": "symmetric_int4_packed",
            "layers": 2,
            "sampling_steps": 8,
            "checkpoint_size_mb": file_size_bytes / (1024 * 1024),
        },
        tags=["1bit", "tensor_core", "int4", "compression", "quantization"],
    )
    tracker.log_benchmark(
        latency_us=int(avg_ms * 1000),
        throughput_items_per_sec=qps,
        device_name=torch.cuda.get_device_name(0),
        notes=f"INT4 outer quantized model: {file_size_bytes / (1024*1024):.2f} MB, {avg_ms:.2f} ms latency",
    )
    tracker.finish(
        status="completed",
        summary_metrics={
            "val_loss": val_loss,
            "prompt_alignment": align,
            "pairwise_diversity": div,
            "latency_ms": avg_ms,
            "file_size_mb": file_size_bytes / (1024 * 1024),
            "file_size_kb": file_size_bytes / 1024,
        },
    )
    print("\nLogged INT4 champion experiment to journal.db!")


if __name__ == "__main__":
    main()