ajh-code commited on
Commit
d178c84
·
verified ·
1 Parent(s): dab3e90

Add runtime/fp4_bridge_runtime.py

Browse files
Files changed (1) hide show
  1. runtime/fp4_bridge_runtime.py +1190 -0
runtime/fp4_bridge_runtime.py ADDED
@@ -0,0 +1,1190 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Portable selected-block FP4 image-MLP bridge runtime for the ComfyUI plugin."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import ctypes
6
+ import math
7
+ from dataclasses import dataclass, field
8
+ from pathlib import Path
9
+ from typing import Any, Iterable, Mapping
10
+
11
+
12
+ RUNTIME_ROOT = Path(__file__).resolve().parent
13
+ ABI_VERSION = 1
14
+ UP_IN_FEATURES = 3072
15
+ UP_OUT_FEATURES = 12288
16
+ DOWN_OUT_FEATURES = 3072
17
+ TRANSFORMER_BLOCK_COUNT = 12
18
+ BLOCK0_IMG_MLP_MODULE = "transformer_blocks.0.img_mlp"
19
+ UP_MODULE = f"{BLOCK0_IMG_MLP_MODULE}.net.0.proj"
20
+ DOWN_MODULE = f"{BLOCK0_IMG_MLP_MODULE}.net.2"
21
+
22
+
23
+ def _tensor_bytes(tensor: Any) -> int:
24
+ return int(tensor.numel() * tensor.element_size())
25
+
26
+
27
+ def resolve_existing_runtime_file(
28
+ path: str | Path,
29
+ label: str,
30
+ *,
31
+ base_dir: str | Path | None = None,
32
+ ) -> Path:
33
+ root = RUNTIME_ROOT if base_dir is None else Path(base_dir).expanduser().resolve()
34
+ candidate = Path(path).expanduser()
35
+ resolved = (root / candidate).resolve() if not candidate.is_absolute() else candidate.resolve()
36
+ if not resolved.is_file():
37
+ raise RuntimeError(f"{label} must resolve to an existing file: {resolved}")
38
+ return resolved
39
+
40
+
41
+ def _require_positive_finite(value: float, label: str) -> float:
42
+ if not math.isfinite(value) or value <= 0.0:
43
+ raise RuntimeError(f"{label} must be finite and positive; got {value!r}")
44
+ return float(value)
45
+
46
+
47
+ def _require_block_index(block_index: int) -> int:
48
+ value = int(block_index)
49
+ if value < 0 or value >= TRANSFORMER_BLOCK_COUNT:
50
+ raise RuntimeError(
51
+ f"bridge block index must be in [0, {TRANSFORMER_BLOCK_COUNT - 1}], got {value}"
52
+ )
53
+ return value
54
+
55
+
56
+ def _module_names_for_block(block_index: int) -> tuple[str, str, str]:
57
+ base = f"transformer_blocks.{block_index}.img_mlp"
58
+ return (
59
+ base,
60
+ f"{base}.net.0.proj",
61
+ f"{base}.net.2",
62
+ )
63
+
64
+
65
+ def normalize_block_tensor_scales(
66
+ block_tensor_scales: str | Mapping[int | str, float | str],
67
+ ) -> dict[int, float]:
68
+ if isinstance(block_tensor_scales, str):
69
+ text = block_tensor_scales.strip()
70
+ if not text:
71
+ raise RuntimeError("bridge block scale map must not be empty")
72
+ parsed: dict[int | str, float | str] = {}
73
+ for entry in text.split(","):
74
+ item = entry.strip()
75
+ if not item:
76
+ raise RuntimeError("bridge block scale map contains an empty entry")
77
+ if "=" not in item:
78
+ raise RuntimeError(
79
+ "bridge block scale map entries must look like block=scale"
80
+ )
81
+ block_text, scale_text = item.split("=", 1)
82
+ block_key = block_text.strip()
83
+ if block_key in parsed:
84
+ raise RuntimeError(f"duplicate bridge block index {block_key}")
85
+ parsed[block_key] = scale_text.strip()
86
+ block_tensor_scales = parsed
87
+ if not block_tensor_scales:
88
+ raise RuntimeError("bridge candidate requires at least one selected block")
89
+ normalized: dict[int, float] = {}
90
+ for raw_block_index, raw_scale in block_tensor_scales.items():
91
+ try:
92
+ block_index = _require_block_index(int(raw_block_index))
93
+ except (TypeError, ValueError) as error:
94
+ raise RuntimeError(
95
+ f"invalid bridge block index {raw_block_index!r}"
96
+ ) from error
97
+ if block_index in normalized:
98
+ raise RuntimeError(f"duplicate bridge block index {block_index}")
99
+ try:
100
+ scale_value = float(raw_scale)
101
+ except (TypeError, ValueError) as error:
102
+ raise RuntimeError(
103
+ f"invalid bridge fixed tensor scale {raw_scale!r} for block {block_index}"
104
+ ) from error
105
+ normalized[block_index] = _require_positive_finite(
106
+ scale_value,
107
+ f"bridge fixed tensor scale for block {block_index}",
108
+ )
109
+ return dict(sorted(normalized.items()))
110
+
111
+
112
+ def normalize_selected_blocks(
113
+ selected_blocks: None | int | str | Iterable[int],
114
+ *,
115
+ allowed_blocks: Iterable[int] | None = None,
116
+ ) -> set[int]:
117
+ if selected_blocks is None:
118
+ return set()
119
+ if isinstance(selected_blocks, int):
120
+ items: list[int | str] = [selected_blocks]
121
+ elif isinstance(selected_blocks, str):
122
+ text = selected_blocks.strip()
123
+ if not text:
124
+ raise RuntimeError("selected_blocks must not be empty")
125
+ items = [item.strip() for item in text.split(",")]
126
+ else:
127
+ items = list(selected_blocks)
128
+ normalized = {_require_block_index(int(item)) for item in items}
129
+ if allowed_blocks is not None:
130
+ allowed = {_require_block_index(int(item)) for item in allowed_blocks}
131
+ unknown = sorted(normalized - allowed)
132
+ if unknown:
133
+ raise RuntimeError(f"selected blocks were not installed: {unknown}")
134
+ return normalized
135
+
136
+
137
+ def _require_packed_projection(
138
+ projection: Any,
139
+ *,
140
+ label: str,
141
+ in_features: int,
142
+ out_features: int,
143
+ torch: Any,
144
+ ) -> None:
145
+ required = (
146
+ "packed_weight",
147
+ "weight_scales",
148
+ "weight_scale",
149
+ "bias",
150
+ "in_features",
151
+ "out_features",
152
+ )
153
+ if projection is None or any(not hasattr(projection, name) for name in required):
154
+ raise RuntimeError(f"{label} is not a packed NVFP4 projection")
155
+ if int(projection.in_features) != in_features:
156
+ raise RuntimeError(
157
+ f"{label} in_features changed: {int(projection.in_features)} != {in_features}"
158
+ )
159
+ if int(projection.out_features) != out_features:
160
+ raise RuntimeError(
161
+ f"{label} out_features changed: {int(projection.out_features)} != {out_features}"
162
+ )
163
+ if projection.bias is None:
164
+ raise RuntimeError(f"{label} requires a resident BF16 bias")
165
+ tensors = (
166
+ ("packed_weight", projection.packed_weight, torch.uint8),
167
+ ("weight_scales", projection.weight_scales, torch.uint8),
168
+ ("weight_scale", projection.weight_scale, torch.float32),
169
+ ("bias", projection.bias, torch.bfloat16),
170
+ )
171
+ for tensor_label, tensor, dtype in tensors:
172
+ if tensor.device.type != "cuda":
173
+ raise RuntimeError(f"{label}.{tensor_label} must be CUDA-resident")
174
+ if tensor.dtype != dtype:
175
+ raise RuntimeError(
176
+ f"{label}.{tensor_label} dtype changed: {tensor.dtype} != {dtype}"
177
+ )
178
+ if not tensor.is_contiguous():
179
+ raise RuntimeError(f"{label}.{tensor_label} must be contiguous")
180
+
181
+
182
+ def _require_exact_numel(
183
+ tensor: Any,
184
+ *,
185
+ expected: int,
186
+ label: str,
187
+ ) -> None:
188
+ if int(tensor.numel()) != int(expected):
189
+ raise RuntimeError(
190
+ f"{label} size changed: {int(tensor.numel())} != {int(expected)}"
191
+ )
192
+
193
+
194
+ def _record_stream_for_tensors(current_stream: Any, *tensors: Any) -> None:
195
+ for tensor in tensors:
196
+ if tensor is not None:
197
+ tensor.record_stream(current_stream)
198
+
199
+
200
+ def _resolve_up_projection(activation: Any, module_name: str) -> tuple[Any, str]:
201
+ projection = getattr(activation, "proj", None)
202
+ if projection is None:
203
+ projection = getattr(activation, "projection", None)
204
+ if projection is None:
205
+ raise RuntimeError(f"{module_name} is missing a .proj/.projection projection")
206
+ if (
207
+ getattr(activation, "is_mage_fused_gelu_up_wrapper", False)
208
+ or getattr(activation, "_xpo3_fused_gelu_up", False)
209
+ ):
210
+ return projection, "fused_gelu_up_wrapper"
211
+ if getattr(activation, "approximate", None) == "tanh":
212
+ return projection, "tanh_gelu"
213
+ raise RuntimeError(
214
+ f"{module_name} activation changed; expected tanh GELU or _xpo3_fused_gelu_up wrapper"
215
+ )
216
+
217
+
218
+ class BridgeUpLibrary:
219
+ def __init__(self, path: Path, device: int):
220
+ self.path = path
221
+ self.library = ctypes.CDLL(str(path), mode=ctypes.RTLD_LOCAL)
222
+ self.context = ctypes.c_void_p()
223
+ self._bind()
224
+ if int(self.library.mage_nvfp4_bridge_up_abi_version()) != ABI_VERSION:
225
+ raise RuntimeError("bridge-up ABI mismatch")
226
+ status = self.library.mage_nvfp4_bridge_up_create_context(
227
+ int(device), ctypes.byref(self.context)
228
+ )
229
+ if status:
230
+ raise self.error("creating bridge-up context")
231
+
232
+ def _bind(self) -> None:
233
+ library = self.library
234
+ library.mage_nvfp4_bridge_up_abi_version.argtypes = []
235
+ library.mage_nvfp4_bridge_up_abi_version.restype = ctypes.c_int
236
+ library.mage_nvfp4_bridge_up_last_error.argtypes = []
237
+ library.mage_nvfp4_bridge_up_last_error.restype = ctypes.c_char_p
238
+ for suffix in (
239
+ "input_bytes",
240
+ "weight_bytes",
241
+ "weight_scale_bytes",
242
+ "output_fp4_bytes",
243
+ "output_scale_bytes",
244
+ ):
245
+ function = getattr(library, f"mage_nvfp4_bridge_up_{suffix}")
246
+ function.argtypes = [ctypes.c_int, ctypes.c_int]
247
+ function.restype = ctypes.c_size_t
248
+ library.mage_nvfp4_bridge_up_bias_bytes.argtypes = [ctypes.c_int]
249
+ library.mage_nvfp4_bridge_up_bias_bytes.restype = ctypes.c_size_t
250
+ library.mage_nvfp4_bridge_up_scalar_bytes.argtypes = []
251
+ library.mage_nvfp4_bridge_up_scalar_bytes.restype = ctypes.c_size_t
252
+ library.mage_nvfp4_bridge_up_create_context.argtypes = [
253
+ ctypes.c_int,
254
+ ctypes.POINTER(ctypes.c_void_p),
255
+ ]
256
+ library.mage_nvfp4_bridge_up_create_context.restype = ctypes.c_int
257
+ library.mage_nvfp4_bridge_up_destroy_context.argtypes = [ctypes.c_void_p]
258
+ library.mage_nvfp4_bridge_up_destroy_context.restype = ctypes.c_int
259
+ library.mage_nvfp4_bridge_up_context_reserved_bytes.argtypes = [
260
+ ctypes.c_void_p
261
+ ]
262
+ library.mage_nvfp4_bridge_up_context_reserved_bytes.restype = ctypes.c_size_t
263
+ forward = library.mage_nvfp4_bridge_up_forward
264
+ forward.argtypes = [
265
+ ctypes.c_void_p,
266
+ ctypes.c_void_p,
267
+ ctypes.c_size_t,
268
+ ctypes.c_void_p,
269
+ ctypes.c_size_t,
270
+ ctypes.c_void_p,
271
+ ctypes.c_size_t,
272
+ ctypes.c_void_p,
273
+ ctypes.c_size_t,
274
+ ctypes.c_void_p,
275
+ ctypes.c_size_t,
276
+ ctypes.c_void_p,
277
+ ctypes.c_size_t,
278
+ ctypes.c_void_p,
279
+ ctypes.c_size_t,
280
+ ctypes.c_void_p,
281
+ ctypes.c_size_t,
282
+ ctypes.c_int,
283
+ ctypes.c_int,
284
+ ctypes.c_int,
285
+ ctypes.c_size_t,
286
+ ]
287
+ forward.restype = ctypes.c_int
288
+ self.forward_function = forward
289
+
290
+ def error(self, operation: str) -> RuntimeError:
291
+ raw = self.library.mage_nvfp4_bridge_up_last_error()
292
+ message = raw.decode("utf-8", errors="replace") if raw else "unknown native error"
293
+ return RuntimeError(f"{operation}: {message}")
294
+
295
+ def helper(self, suffix: str, *dimensions: int) -> int:
296
+ function = getattr(self.library, f"mage_nvfp4_bridge_up_{suffix}")
297
+ value = int(function(*dimensions))
298
+ if value == 0:
299
+ raise self.error(f"querying bridge-up {suffix}")
300
+ return value
301
+
302
+ def scalar_bytes(self) -> int:
303
+ value = int(self.library.mage_nvfp4_bridge_up_scalar_bytes())
304
+ if value == 0:
305
+ raise self.error("querying bridge-up scalar bytes")
306
+ return value
307
+
308
+ @property
309
+ def reserved_bytes(self) -> int:
310
+ return int(
311
+ self.library.mage_nvfp4_bridge_up_context_reserved_bytes(self.context)
312
+ )
313
+
314
+ def forward(
315
+ self,
316
+ input_bf16: Any,
317
+ packed_weight: Any,
318
+ packed_weight_scales: Any,
319
+ weight_tensor_scale: Any,
320
+ bias_bf16: Any,
321
+ bridge_norm_constant: Any,
322
+ output_fp4: Any,
323
+ output_scales: Any,
324
+ logical_m: int,
325
+ stream: int,
326
+ ) -> None:
327
+ status = self.forward_function(
328
+ self.context,
329
+ ctypes.c_void_p(input_bf16.data_ptr()),
330
+ _tensor_bytes(input_bf16),
331
+ ctypes.c_void_p(packed_weight.data_ptr()),
332
+ _tensor_bytes(packed_weight),
333
+ ctypes.c_void_p(packed_weight_scales.data_ptr()),
334
+ _tensor_bytes(packed_weight_scales),
335
+ ctypes.c_void_p(weight_tensor_scale.data_ptr()),
336
+ _tensor_bytes(weight_tensor_scale),
337
+ ctypes.c_void_p(bias_bf16.data_ptr()),
338
+ _tensor_bytes(bias_bf16),
339
+ ctypes.c_void_p(bridge_norm_constant.data_ptr()),
340
+ _tensor_bytes(bridge_norm_constant),
341
+ ctypes.c_void_p(output_fp4.data_ptr()),
342
+ _tensor_bytes(output_fp4),
343
+ ctypes.c_void_p(output_scales.data_ptr()),
344
+ _tensor_bytes(output_scales),
345
+ int(logical_m),
346
+ UP_IN_FEATURES,
347
+ UP_OUT_FEATURES,
348
+ int(stream),
349
+ )
350
+ if status:
351
+ raise self.error("running bridge-up")
352
+
353
+ def close(self) -> None:
354
+ if not self.context:
355
+ return
356
+ status = self.library.mage_nvfp4_bridge_up_destroy_context(self.context)
357
+ self.context = ctypes.c_void_p()
358
+ if status:
359
+ raise self.error("destroying bridge-up context")
360
+
361
+
362
+ class PrequantizedDownLibrary:
363
+ def __init__(self, path: Path, device: int):
364
+ self.path = path
365
+ self.library = ctypes.CDLL(str(path), mode=ctypes.RTLD_LOCAL)
366
+ self.context = ctypes.c_void_p()
367
+ self._bind()
368
+ if (
369
+ int(self.library.mage_nvfp4_prequantized_down_abi_version())
370
+ != ABI_VERSION
371
+ ):
372
+ raise RuntimeError("bridge-down ABI mismatch")
373
+ status = self.library.mage_nvfp4_prequantized_down_create_context(
374
+ int(device), ctypes.byref(self.context)
375
+ )
376
+ if status:
377
+ raise self.error("creating bridge-down context")
378
+
379
+ def _bind(self) -> None:
380
+ library = self.library
381
+ prefix = "mage_nvfp4_prequantized_down_"
382
+ abi = getattr(library, f"{prefix}abi_version")
383
+ abi.argtypes = []
384
+ abi.restype = ctypes.c_int
385
+ last_error = getattr(library, f"{prefix}last_error")
386
+ last_error.argtypes = []
387
+ last_error.restype = ctypes.c_char_p
388
+ for suffix in (
389
+ "activation_bytes",
390
+ "activation_scale_bytes",
391
+ "weight_bytes",
392
+ "weight_scale_bytes",
393
+ "output_bytes",
394
+ ):
395
+ function = getattr(library, f"{prefix}{suffix}")
396
+ function.argtypes = [ctypes.c_int, ctypes.c_int]
397
+ function.restype = ctypes.c_size_t
398
+ bias_bytes = getattr(library, f"{prefix}bias_bytes")
399
+ bias_bytes.argtypes = [ctypes.c_int]
400
+ bias_bytes.restype = ctypes.c_size_t
401
+ create = getattr(library, f"{prefix}create_context")
402
+ create.argtypes = [ctypes.c_int, ctypes.POINTER(ctypes.c_void_p)]
403
+ create.restype = ctypes.c_int
404
+ destroy = getattr(library, f"{prefix}destroy_context")
405
+ destroy.argtypes = [ctypes.c_void_p]
406
+ destroy.restype = ctypes.c_int
407
+ reserved = getattr(library, f"{prefix}context_reserved_bytes")
408
+ reserved.argtypes = [ctypes.c_void_p]
409
+ reserved.restype = ctypes.c_size_t
410
+ forward = getattr(library, f"{prefix}forward")
411
+ forward.argtypes = [
412
+ ctypes.c_void_p,
413
+ ctypes.c_void_p,
414
+ ctypes.c_size_t,
415
+ ctypes.c_void_p,
416
+ ctypes.c_size_t,
417
+ ctypes.c_void_p,
418
+ ctypes.c_size_t,
419
+ ctypes.c_void_p,
420
+ ctypes.c_size_t,
421
+ ctypes.c_void_p,
422
+ ctypes.c_size_t,
423
+ ctypes.c_void_p,
424
+ ctypes.c_size_t,
425
+ ctypes.c_void_p,
426
+ ctypes.c_size_t,
427
+ ctypes.c_void_p,
428
+ ctypes.c_size_t,
429
+ ctypes.c_int,
430
+ ctypes.c_int,
431
+ ctypes.c_int,
432
+ ctypes.c_size_t,
433
+ ]
434
+ forward.restype = ctypes.c_int
435
+ self.forward_function = forward
436
+
437
+ def error(self, operation: str) -> RuntimeError:
438
+ raw = self.library.mage_nvfp4_prequantized_down_last_error()
439
+ message = raw.decode("utf-8", errors="replace") if raw else "unknown native error"
440
+ return RuntimeError(f"{operation}: {message}")
441
+
442
+ def helper(self, suffix: str, *dimensions: int) -> int:
443
+ function = getattr(self.library, f"mage_nvfp4_prequantized_down_{suffix}")
444
+ value = int(function(*dimensions))
445
+ if value == 0:
446
+ raise self.error(f"querying bridge-down {suffix}")
447
+ return value
448
+
449
+ @property
450
+ def reserved_bytes(self) -> int:
451
+ return int(
452
+ self.library.mage_nvfp4_prequantized_down_context_reserved_bytes(
453
+ self.context
454
+ )
455
+ )
456
+
457
+ def forward(
458
+ self,
459
+ activation_fp4: Any,
460
+ activation_scales: Any,
461
+ activation_tensor_scale: Any,
462
+ packed_weight: Any,
463
+ packed_weight_scales: Any,
464
+ weight_tensor_scale: Any,
465
+ bias_bf16: Any,
466
+ output_bf16: Any,
467
+ logical_m: int,
468
+ stream: int,
469
+ ) -> None:
470
+ status = self.forward_function(
471
+ self.context,
472
+ ctypes.c_void_p(activation_fp4.data_ptr()),
473
+ _tensor_bytes(activation_fp4),
474
+ ctypes.c_void_p(activation_scales.data_ptr()),
475
+ _tensor_bytes(activation_scales),
476
+ ctypes.c_void_p(activation_tensor_scale.data_ptr()),
477
+ _tensor_bytes(activation_tensor_scale),
478
+ ctypes.c_void_p(packed_weight.data_ptr()),
479
+ _tensor_bytes(packed_weight),
480
+ ctypes.c_void_p(packed_weight_scales.data_ptr()),
481
+ _tensor_bytes(packed_weight_scales),
482
+ ctypes.c_void_p(weight_tensor_scale.data_ptr()),
483
+ _tensor_bytes(weight_tensor_scale),
484
+ ctypes.c_void_p(bias_bf16.data_ptr()),
485
+ _tensor_bytes(bias_bf16),
486
+ ctypes.c_void_p(output_bf16.data_ptr()),
487
+ _tensor_bytes(output_bf16),
488
+ int(logical_m),
489
+ UP_OUT_FEATURES,
490
+ DOWN_OUT_FEATURES,
491
+ int(stream),
492
+ )
493
+ if status:
494
+ raise self.error("running bridge-down")
495
+
496
+ def close(self) -> None:
497
+ if not self.context:
498
+ return
499
+ status = self.library.mage_nvfp4_prequantized_down_destroy_context(
500
+ self.context
501
+ )
502
+ self.context = ctypes.c_void_p()
503
+ if status:
504
+ raise self.error("destroying bridge-down context")
505
+
506
+
507
+ @dataclass
508
+ class _BridgeBuffers:
509
+ payload: Any
510
+ scales: Any
511
+ output: Any
512
+
513
+
514
+ @dataclass(frozen=True)
515
+ class _BridgeBlockBinding:
516
+ block_index: int
517
+ module: str
518
+ up_module: str
519
+ down_module: str
520
+ up_projection: Any
521
+ down_projection: Any
522
+ fixed_tensor_scale: float
523
+ bridge_tensor_scale: Any
524
+ bridge_norm_constant: Any
525
+ activation_mode: str
526
+
527
+
528
+ @dataclass
529
+ class _BlockTelemetry:
530
+ native_calls: int = 0
531
+ fallback_calls: int = 0
532
+ fallback_reasons: dict[str, int] = field(default_factory=dict)
533
+ last_route: str | None = None
534
+ last_logical_m: int | None = None
535
+ last_stream: int | None = None
536
+
537
+ def record(self, *, native: bool, reason: str, logical_m: int | None, stream: int | None) -> None:
538
+ self.last_route = reason
539
+ self.last_logical_m = logical_m
540
+ self.last_stream = stream
541
+ if native:
542
+ self.native_calls += 1
543
+ return
544
+ self.fallback_calls += 1
545
+ self.fallback_reasons[reason] = self.fallback_reasons.get(reason, 0) + 1
546
+
547
+ def snapshot(self) -> dict[str, Any]:
548
+ return {
549
+ "native_calls": self.native_calls,
550
+ "fallback_calls": self.fallback_calls,
551
+ "fallback_reasons": dict(sorted(self.fallback_reasons.items())),
552
+ "last_route": self.last_route,
553
+ "last_logical_m": self.last_logical_m,
554
+ "last_stream": self.last_stream,
555
+ }
556
+
557
+
558
+ class BlockImgMlpBridgeRuntime:
559
+ def __init__(
560
+ self,
561
+ *,
562
+ bridge_up_library_path: Path,
563
+ bridge_down_library_path: Path,
564
+ torch: Any,
565
+ device_index: int | None = None,
566
+ bridge_up_factory: Any = BridgeUpLibrary,
567
+ bridge_down_factory: Any = PrequantizedDownLibrary,
568
+ ) -> None:
569
+ self.torch = torch
570
+ self.bridge_up_library_path = Path(bridge_up_library_path).resolve()
571
+ self.bridge_down_library_path = Path(bridge_down_library_path).resolve()
572
+ self.device_index = (
573
+ int(device_index)
574
+ if device_index is not None
575
+ else int(torch.cuda.current_device())
576
+ )
577
+ self.bridge_up = bridge_up_factory(
578
+ self.bridge_up_library_path,
579
+ self.device_index,
580
+ )
581
+ self.bridge_down = bridge_down_factory(
582
+ self.bridge_down_library_path,
583
+ self.device_index,
584
+ )
585
+ scalar_bytes = self.bridge_up.scalar_bytes()
586
+ if scalar_bytes != 4:
587
+ raise RuntimeError(
588
+ f"bridge scalar ABI changed: expected 4 bytes, got {scalar_bytes}"
589
+ )
590
+ self._buffers_by_m: dict[int, _BridgeBuffers] = {}
591
+ self._bindings_by_block: dict[int, _BridgeBlockBinding] = {}
592
+ self._telemetry_by_block: dict[int, _BlockTelemetry] = {}
593
+ self._enabled = True
594
+ self._enabled_blocks: set[int] = set()
595
+ self._bound_stream: int | None = None
596
+ self._closed = False
597
+
598
+ def _telemetry(self, block_index: int) -> _BlockTelemetry:
599
+ return self._telemetry_by_block.setdefault(block_index, _BlockTelemetry())
600
+
601
+ def _logical_m_reason(self, logical_m: int) -> str | None:
602
+ if logical_m <= 0:
603
+ return "logical_m_non_positive"
604
+ if logical_m % 8:
605
+ return "logical_m_not_divisible_by_8"
606
+ if logical_m > 6400:
607
+ return "logical_m_above_max"
608
+ return None
609
+
610
+ def _require_logical_m(self, logical_m: int) -> None:
611
+ reason = self._logical_m_reason(logical_m)
612
+ if reason == "logical_m_non_positive":
613
+ raise RuntimeError("bridge candidate requires a positive logical M")
614
+ if reason == "logical_m_not_divisible_by_8":
615
+ raise RuntimeError(
616
+ "bridge candidate only supports logical M divisible by 8"
617
+ )
618
+ if reason == "logical_m_above_max":
619
+ raise RuntimeError(
620
+ "bridge candidate only supports logical M up to 6400"
621
+ )
622
+
623
+ def _buffers(self, logical_m: int, device: Any) -> _BridgeBuffers:
624
+ self._require_logical_m(logical_m)
625
+ cached = self._buffers_by_m.get(int(logical_m))
626
+ if cached is not None:
627
+ return cached
628
+ payload_bytes = self.bridge_up.helper("output_fp4_bytes", logical_m, UP_OUT_FEATURES)
629
+ scale_bytes = self.bridge_up.helper("output_scale_bytes", logical_m, UP_OUT_FEATURES)
630
+ output_bytes = self.bridge_down.helper("output_bytes", logical_m, DOWN_OUT_FEATURES)
631
+ output_elements = output_bytes // self.torch.tensor(
632
+ [],
633
+ dtype=self.torch.bfloat16,
634
+ ).element_size()
635
+ expected_elements = int(logical_m) * DOWN_OUT_FEATURES
636
+ if output_elements != expected_elements:
637
+ raise RuntimeError(
638
+ "bridge-down output byte contract changed: "
639
+ f"{output_elements} elements != {expected_elements}"
640
+ )
641
+ buffers = _BridgeBuffers(
642
+ payload=self.torch.empty(payload_bytes, dtype=self.torch.uint8, device=device),
643
+ scales=self.torch.empty(scale_bytes, dtype=self.torch.uint8, device=device),
644
+ output=self.torch.empty(
645
+ expected_elements,
646
+ dtype=self.torch.bfloat16,
647
+ device=device,
648
+ ).view(int(logical_m), DOWN_OUT_FEATURES),
649
+ )
650
+ self._buffers_by_m[int(logical_m)] = buffers
651
+ return buffers
652
+
653
+ def bind_block(
654
+ self,
655
+ *,
656
+ block_index: int,
657
+ module: str,
658
+ up_module: str,
659
+ down_module: str,
660
+ up_projection: Any,
661
+ down_projection: Any,
662
+ fixed_tensor_scale: float,
663
+ activation_mode: str,
664
+ ) -> _BridgeBlockBinding:
665
+ checked_block_index = _require_block_index(block_index)
666
+ if checked_block_index in self._bindings_by_block:
667
+ raise RuntimeError(f"bridge block {checked_block_index} already installed")
668
+ scale_value = _require_positive_finite(
669
+ fixed_tensor_scale,
670
+ f"bridge fixed tensor scale for block {checked_block_index}",
671
+ )
672
+ binding = _BridgeBlockBinding(
673
+ block_index=checked_block_index,
674
+ module=module,
675
+ up_module=up_module,
676
+ down_module=down_module,
677
+ up_projection=up_projection,
678
+ down_projection=down_projection,
679
+ fixed_tensor_scale=scale_value,
680
+ bridge_tensor_scale=self.torch.tensor(
681
+ [scale_value],
682
+ device=f"cuda:{self.device_index}",
683
+ dtype=self.torch.float32,
684
+ ),
685
+ bridge_norm_constant=self.torch.tensor(
686
+ [1.0 / scale_value],
687
+ device=f"cuda:{self.device_index}",
688
+ dtype=self.torch.float32,
689
+ ),
690
+ activation_mode=activation_mode,
691
+ )
692
+ self._bindings_by_block[checked_block_index] = binding
693
+ self._enabled_blocks.add(checked_block_index)
694
+ self._telemetry(checked_block_index)
695
+ return binding
696
+
697
+ def installed_block_indices(self) -> list[int]:
698
+ return sorted(self._bindings_by_block)
699
+
700
+ @property
701
+ def enabled(self) -> bool:
702
+ return self._enabled and not self._closed
703
+
704
+ @property
705
+ def enabled_block_indices(self) -> list[int]:
706
+ return sorted(self._enabled_blocks)
707
+
708
+ def set_enabled(
709
+ self,
710
+ enabled: bool,
711
+ selected_blocks: None | int | str | Iterable[int] = None,
712
+ ) -> None:
713
+ if enabled and self._closed:
714
+ raise RuntimeError("FP4 bridge runtime is closed")
715
+ if selected_blocks is None:
716
+ self._enabled = bool(enabled)
717
+ return
718
+ blocks = normalize_selected_blocks(
719
+ selected_blocks,
720
+ allowed_blocks=self._bindings_by_block,
721
+ )
722
+ if enabled:
723
+ self._enabled_blocks.update(blocks)
724
+ else:
725
+ self._enabled_blocks.difference_update(blocks)
726
+
727
+ def set_active_blocks(
728
+ self,
729
+ selected_blocks: None | int | str | Iterable[int],
730
+ ) -> None:
731
+ if self._closed:
732
+ raise RuntimeError("FP4 bridge runtime is closed")
733
+ self._enabled_blocks = normalize_selected_blocks(
734
+ selected_blocks,
735
+ allowed_blocks=self._bindings_by_block,
736
+ )
737
+
738
+ def reset_telemetry(self) -> None:
739
+ self._telemetry_by_block = {
740
+ block_index: _BlockTelemetry()
741
+ for block_index in self._bindings_by_block
742
+ }
743
+
744
+ def _logical_m_from_shape(self, hidden_states: Any) -> int | None:
745
+ shape = getattr(hidden_states, "shape", None)
746
+ if shape is None:
747
+ return None
748
+ if len(shape) < 1:
749
+ return None
750
+ return int(math.prod(int(value) for value in shape[:-1])) if len(shape) > 1 else 1
751
+
752
+ def _current_stream_handle(self, hidden_states: Any) -> int:
753
+ current_stream = self.torch.cuda.current_stream(hidden_states.device)
754
+ return int(current_stream.cuda_stream)
755
+
756
+ def _route_reason(
757
+ self,
758
+ hidden_states: Any,
759
+ binding: _BridgeBlockBinding,
760
+ ) -> tuple[bool, str, int | None, int | None]:
761
+ logical_m = self._logical_m_from_shape(hidden_states)
762
+ if self._closed:
763
+ return False, "runtime_closed", logical_m, None
764
+ if not self._enabled:
765
+ return False, "runtime_disabled", logical_m, None
766
+ if binding.block_index not in self._enabled_blocks:
767
+ return False, "block_disabled", logical_m, None
768
+ if getattr(hidden_states, "device", None) is None:
769
+ return False, "missing_device", logical_m, None
770
+ if hidden_states.device.type != "cuda":
771
+ return False, f"device_{hidden_states.device.type}", logical_m, None
772
+ if getattr(hidden_states, "dtype", None) != self.torch.bfloat16:
773
+ return False, f"dtype_{hidden_states.dtype}", logical_m, None
774
+ if getattr(hidden_states, "ndim", 0) < 1:
775
+ return False, "shape_rank", logical_m, None
776
+ if int(hidden_states.shape[-1]) != UP_IN_FEATURES:
777
+ return False, "shape_last_dim", logical_m, None
778
+ if logical_m is None:
779
+ return False, "logical_m_unknown", logical_m, None
780
+ logical_reason = self._logical_m_reason(int(logical_m))
781
+ if logical_reason is not None:
782
+ return False, logical_reason, int(logical_m), None
783
+ try:
784
+ stream = self._current_stream_handle(hidden_states)
785
+ except Exception:
786
+ return False, "stream_query_failed", int(logical_m), None
787
+ if self._bound_stream is not None and self._bound_stream != stream:
788
+ return False, "stream_mismatch", int(logical_m), stream
789
+ return True, "native", int(logical_m), stream
790
+
791
+ def _native_forward(
792
+ self,
793
+ hidden_states: Any,
794
+ binding: _BridgeBlockBinding,
795
+ *,
796
+ logical_m: int,
797
+ stream: int,
798
+ ) -> Any:
799
+ flattened = hidden_states.reshape(-1, UP_IN_FEATURES).contiguous()
800
+ buffers = self._buffers(logical_m, flattened.device)
801
+ current_stream = self.torch.cuda.current_stream(hidden_states.device)
802
+ if self._bound_stream is None:
803
+ self._bound_stream = stream
804
+ self.bridge_up.forward(
805
+ flattened,
806
+ binding.up_projection.packed_weight,
807
+ binding.up_projection.weight_scales,
808
+ binding.up_projection.weight_scale,
809
+ binding.up_projection.bias,
810
+ binding.bridge_norm_constant,
811
+ buffers.payload,
812
+ buffers.scales,
813
+ logical_m,
814
+ stream,
815
+ )
816
+ self.bridge_down.forward(
817
+ buffers.payload,
818
+ buffers.scales,
819
+ binding.bridge_tensor_scale,
820
+ binding.down_projection.packed_weight,
821
+ binding.down_projection.weight_scales,
822
+ binding.down_projection.weight_scale,
823
+ binding.down_projection.bias,
824
+ buffers.output,
825
+ logical_m,
826
+ stream,
827
+ )
828
+ _record_stream_for_tensors(
829
+ current_stream,
830
+ flattened,
831
+ binding.up_projection.packed_weight,
832
+ binding.up_projection.weight_scales,
833
+ binding.up_projection.weight_scale,
834
+ binding.up_projection.bias,
835
+ binding.bridge_norm_constant,
836
+ buffers.payload,
837
+ buffers.scales,
838
+ binding.bridge_tensor_scale,
839
+ binding.down_projection.packed_weight,
840
+ binding.down_projection.weight_scales,
841
+ binding.down_projection.weight_scale,
842
+ binding.down_projection.bias,
843
+ buffers.output,
844
+ )
845
+ return buffers.output.view(*hidden_states.shape[:-1], DOWN_OUT_FEATURES)
846
+
847
+ def forward_or_fallback(
848
+ self,
849
+ hidden_states: Any,
850
+ binding: _BridgeBlockBinding,
851
+ fallback_module: Any,
852
+ ) -> Any:
853
+ can_launch, reason, logical_m, stream = self._route_reason(
854
+ hidden_states,
855
+ binding,
856
+ )
857
+ telemetry = self._telemetry(binding.block_index)
858
+ if not can_launch:
859
+ telemetry.record(
860
+ native=False,
861
+ reason=reason,
862
+ logical_m=logical_m,
863
+ stream=stream,
864
+ )
865
+ return fallback_module(hidden_states)
866
+ try:
867
+ output = self._native_forward(
868
+ hidden_states,
869
+ binding,
870
+ logical_m=int(logical_m),
871
+ stream=int(stream),
872
+ )
873
+ except Exception:
874
+ telemetry.record(
875
+ native=False,
876
+ reason="native_error",
877
+ logical_m=logical_m,
878
+ stream=stream,
879
+ )
880
+ # A native error can occur after work was enqueued. Do not mix an
881
+ # ordinary fallback launch into that same stream.
882
+ raise
883
+ telemetry.record(
884
+ native=True,
885
+ reason="native",
886
+ logical_m=logical_m,
887
+ stream=stream,
888
+ )
889
+ return output
890
+
891
+ def telemetry_snapshot(self) -> dict[str, Any]:
892
+ return {
893
+ str(block_index): self._telemetry(block_index).snapshot()
894
+ for block_index in sorted(self._bindings_by_block)
895
+ }
896
+
897
+ def report(self) -> dict[str, Any]:
898
+ return {
899
+ "bridge_up_library": str(self.bridge_up_library_path),
900
+ "bridge_down_library": str(self.bridge_down_library_path),
901
+ "block_indices": self.installed_block_indices(),
902
+ "enabled": self.enabled,
903
+ "enabled_block_indices": self.enabled_block_indices,
904
+ "logical_m_contract": {
905
+ "divisible_by_8": True,
906
+ "max_supported": 6400,
907
+ "fallback_on_unsupported": True,
908
+ },
909
+ "toggle_contract": {
910
+ "default_enabled": True,
911
+ "global_gate": "set_enabled(bool)",
912
+ "block_gate": "set_enabled(bool, selected_blocks=...)",
913
+ "native_route_requires_global_and_block_enable": True,
914
+ },
915
+ "telemetry_by_block": self.telemetry_snapshot(),
916
+ "closed": self._closed,
917
+ }
918
+
919
+ def close(self) -> None:
920
+ if self._closed:
921
+ return
922
+ self._closed = True
923
+ self._enabled = False
924
+ if self.torch.cuda.is_available():
925
+ self.torch.cuda.synchronize(self.device_index)
926
+ close_error = None
927
+ for library in (self.bridge_down, self.bridge_up):
928
+ try:
929
+ library.close()
930
+ except Exception as error: # noqa: BLE001
931
+ if close_error is None:
932
+ close_error = error
933
+ self._bindings_by_block.clear()
934
+ self._enabled_blocks.clear()
935
+ self._buffers_by_m.clear()
936
+ if close_error is not None:
937
+ raise close_error
938
+
939
+
940
+ def install_selected_img_mlp_bridges(
941
+ transformer: Any,
942
+ *,
943
+ bridge_up_library_path: str | Path,
944
+ bridge_down_library_path: str | Path,
945
+ block_tensor_scales: str | Mapping[int | str, float | str],
946
+ torch: Any,
947
+ enabled: bool = True,
948
+ device_index: int | None = None,
949
+ bridge_up_factory: Any = BridgeUpLibrary,
950
+ bridge_down_factory: Any = PrequantizedDownLibrary,
951
+ ) -> tuple[BlockImgMlpBridgeRuntime, dict[str, Any]]:
952
+ """Replace selected image MLP blocks with the chained FP4 bridge pair."""
953
+
954
+ import torch.nn as nn
955
+
956
+ if not torch.cuda.is_available():
957
+ raise RuntimeError("bridge candidate installation requires CUDA")
958
+ if transformer.training:
959
+ raise RuntimeError("bridge candidate expects an eval-mode transformer")
960
+ normalized_scales = normalize_block_tensor_scales(block_tensor_scales)
961
+ runtime = BlockImgMlpBridgeRuntime(
962
+ bridge_up_library_path=resolve_existing_runtime_file(
963
+ bridge_up_library_path,
964
+ "bridge-up library",
965
+ ),
966
+ bridge_down_library_path=resolve_existing_runtime_file(
967
+ bridge_down_library_path,
968
+ "bridge-down library",
969
+ ),
970
+ torch=torch,
971
+ device_index=device_index,
972
+ bridge_up_factory=bridge_up_factory,
973
+ bridge_down_factory=bridge_down_factory,
974
+ )
975
+ installed: list[dict[str, Any]] = []
976
+ try:
977
+ for block_index in sorted(normalized_scales):
978
+ module_name, up_module_name, down_module_name = _module_names_for_block(
979
+ block_index
980
+ )
981
+ block = transformer.transformer_blocks[block_index]
982
+ feed_forward = block.img_mlp
983
+ net = getattr(feed_forward, "net", None)
984
+ if net is None or len(net) != 3:
985
+ raise RuntimeError(
986
+ f"{module_name} layout changed; expected 3 net entries"
987
+ )
988
+ activation = net[0]
989
+ up_projection, activation_mode = _resolve_up_projection(
990
+ activation,
991
+ up_module_name,
992
+ )
993
+ dropout = net[1]
994
+ if not isinstance(dropout, nn.Dropout):
995
+ raise RuntimeError(f"{module_name}.net.1 changed; expected Dropout")
996
+ down_projection = net[2]
997
+ try:
998
+ _require_packed_projection(
999
+ up_projection,
1000
+ label=up_module_name,
1001
+ in_features=UP_IN_FEATURES,
1002
+ out_features=UP_OUT_FEATURES,
1003
+ torch=torch,
1004
+ )
1005
+ _require_packed_projection(
1006
+ down_projection,
1007
+ label=down_module_name,
1008
+ in_features=UP_OUT_FEATURES,
1009
+ out_features=DOWN_OUT_FEATURES,
1010
+ torch=torch,
1011
+ )
1012
+ except RuntimeError as error:
1013
+ raise RuntimeError(
1014
+ f"{module_name} cannot use the bridge candidate because it is not a packed NVFP4 image MLP block: {error}"
1015
+ ) from error
1016
+ if up_projection.packed_weight.device != down_projection.packed_weight.device:
1017
+ raise RuntimeError(
1018
+ f"{module_name} requires up/down projections on one device"
1019
+ )
1020
+ _require_exact_numel(
1021
+ up_projection.packed_weight,
1022
+ expected=runtime.bridge_up.helper(
1023
+ "weight_bytes",
1024
+ UP_OUT_FEATURES,
1025
+ UP_IN_FEATURES,
1026
+ ),
1027
+ label=f"{up_module_name}.packed_weight",
1028
+ )
1029
+ _require_exact_numel(
1030
+ up_projection.weight_scales,
1031
+ expected=runtime.bridge_up.helper(
1032
+ "weight_scale_bytes",
1033
+ UP_OUT_FEATURES,
1034
+ UP_IN_FEATURES,
1035
+ ),
1036
+ label=f"{up_module_name}.weight_scales",
1037
+ )
1038
+ _require_exact_numel(
1039
+ up_projection.bias,
1040
+ expected=runtime.bridge_up.helper("bias_bytes", UP_OUT_FEATURES)
1041
+ // up_projection.bias.element_size(),
1042
+ label=f"{up_module_name}.bias",
1043
+ )
1044
+ _require_exact_numel(
1045
+ down_projection.packed_weight,
1046
+ expected=runtime.bridge_down.helper(
1047
+ "weight_bytes",
1048
+ DOWN_OUT_FEATURES,
1049
+ UP_OUT_FEATURES,
1050
+ ),
1051
+ label=f"{down_module_name}.packed_weight",
1052
+ )
1053
+ _require_exact_numel(
1054
+ down_projection.weight_scales,
1055
+ expected=runtime.bridge_down.helper(
1056
+ "weight_scale_bytes",
1057
+ DOWN_OUT_FEATURES,
1058
+ UP_OUT_FEATURES,
1059
+ ),
1060
+ label=f"{down_module_name}.weight_scales",
1061
+ )
1062
+ _require_exact_numel(
1063
+ down_projection.bias,
1064
+ expected=runtime.bridge_down.helper("bias_bytes", DOWN_OUT_FEATURES)
1065
+ // down_projection.bias.element_size(),
1066
+ label=f"{down_module_name}.bias",
1067
+ )
1068
+ binding = runtime.bind_block(
1069
+ block_index=block_index,
1070
+ module=module_name,
1071
+ up_module=up_module_name,
1072
+ down_module=down_module_name,
1073
+ up_projection=up_projection,
1074
+ down_projection=down_projection,
1075
+ fixed_tensor_scale=normalized_scales[block_index],
1076
+ activation_mode=activation_mode,
1077
+ )
1078
+
1079
+ class BlockImgMlpBridge(nn.Module):
1080
+ def __init__(
1081
+ self,
1082
+ installed_binding: _BridgeBlockBinding,
1083
+ original_feed_forward: Any,
1084
+ ) -> None:
1085
+ super().__init__()
1086
+ self.binding = installed_binding
1087
+ self.fallback_module = original_feed_forward
1088
+ self._xpo3_fp4_bridge = True
1089
+ self._xpo3_fp4_bridge_block_index = int(
1090
+ installed_binding.block_index
1091
+ )
1092
+
1093
+ def forward(self, hidden_states: Any) -> Any:
1094
+ return runtime.forward_or_fallback(
1095
+ hidden_states,
1096
+ self.binding,
1097
+ self.fallback_module,
1098
+ )
1099
+
1100
+ block.img_mlp = BlockImgMlpBridge(binding, feed_forward)
1101
+ installed.append(
1102
+ {
1103
+ "block_index": block_index,
1104
+ "module": module_name,
1105
+ "up_module": up_module_name,
1106
+ "down_module": down_module_name,
1107
+ "fixed_tensor_scale": float(binding.fixed_tensor_scale),
1108
+ "bridge_norm_constant": float(binding.bridge_norm_constant.item()),
1109
+ "activation_mode": activation_mode,
1110
+ }
1111
+ )
1112
+ runtime.set_enabled(bool(enabled))
1113
+ except Exception:
1114
+ runtime.close()
1115
+ raise
1116
+ metadata = runtime.report()
1117
+ metadata.update(
1118
+ {
1119
+ "mode": "selected_img_mlp",
1120
+ "block_count": len(installed),
1121
+ "blocks": installed,
1122
+ "fixed_tensor_scales_by_block": {
1123
+ str(entry["block_index"]): entry["fixed_tensor_scale"]
1124
+ for entry in installed
1125
+ },
1126
+ "bridge_up_context_reserved_bytes": runtime.bridge_up.reserved_bytes,
1127
+ "bridge_down_context_reserved_bytes": runtime.bridge_down.reserved_bytes,
1128
+ "resident_scalar_bytes": (
1129
+ sum(
1130
+ _tensor_bytes(binding.bridge_tensor_scale)
1131
+ for binding in runtime._bindings_by_block.values()
1132
+ )
1133
+ + sum(
1134
+ _tensor_bytes(binding.bridge_norm_constant)
1135
+ for binding in runtime._bindings_by_block.values()
1136
+ )
1137
+ ),
1138
+ "reuses_existing_projection_buffers": True,
1139
+ "validation": {
1140
+ "no_bf16_post_gelu_intermediate_materialized": True,
1141
+ "bf16_input_required": True,
1142
+ "bf16_output_preserved": True,
1143
+ "dropout_bypassed_in_eval_only": True,
1144
+ "single_cuda_stream_required_for_native_route": True,
1145
+ "unsupported_conditions_fallback_before_native_launch": True,
1146
+ },
1147
+ }
1148
+ )
1149
+ return runtime, metadata
1150
+
1151
+
1152
+ def install_block0_img_mlp_bridge(
1153
+ transformer: Any,
1154
+ *,
1155
+ bridge_up_library_path: str | Path,
1156
+ bridge_down_library_path: str | Path,
1157
+ fixed_tensor_scale: float,
1158
+ torch: Any,
1159
+ enabled: bool = True,
1160
+ device_index: int | None = None,
1161
+ bridge_up_factory: Any = BridgeUpLibrary,
1162
+ bridge_down_factory: Any = PrequantizedDownLibrary,
1163
+ ) -> tuple[BlockImgMlpBridgeRuntime, dict[str, Any]]:
1164
+ runtime, metadata = install_selected_img_mlp_bridges(
1165
+ transformer,
1166
+ bridge_up_library_path=bridge_up_library_path,
1167
+ bridge_down_library_path=bridge_down_library_path,
1168
+ block_tensor_scales={0: fixed_tensor_scale},
1169
+ torch=torch,
1170
+ enabled=enabled,
1171
+ device_index=device_index,
1172
+ bridge_up_factory=bridge_up_factory,
1173
+ bridge_down_factory=bridge_down_factory,
1174
+ )
1175
+ metadata["mode"] = "block0_img_mlp"
1176
+ return runtime, metadata
1177
+
1178
+
1179
+ __all__ = [
1180
+ "BLOCK0_IMG_MLP_MODULE",
1181
+ "DOWN_MODULE",
1182
+ "TRANSFORMER_BLOCK_COUNT",
1183
+ "UP_MODULE",
1184
+ "BlockImgMlpBridgeRuntime",
1185
+ "install_block0_img_mlp_bridge",
1186
+ "install_selected_img_mlp_bridges",
1187
+ "normalize_block_tensor_scales",
1188
+ "normalize_selected_blocks",
1189
+ "resolve_existing_runtime_file",
1190
+ ]