kernels-bot commited on
Commit
eaedc44
·
verified ·
1 Parent(s): 811eb42

Uploaded using `kernel-builder`.

Browse files
build/torch-rocm/_ops.py CHANGED
@@ -22,7 +22,7 @@ def get_backend() -> str:
22
 
23
  def _find_ops_name() -> str:
24
  kernel_name = "finegrained_kernels"
25
- unique_id = "1348465"
26
  backend = get_backend()
27
  return f"_{kernel_name}_{backend}_{unique_id}"
28
 
 
22
 
23
  def _find_ops_name() -> str:
24
  kernel_name = "finegrained_kernels"
25
+ unique_id = "b0c6080"
26
  backend = get_backend()
27
  return f"_{kernel_name}_{backend}_{unique_id}"
28
 
build/torch-rocm/batched.py CHANGED
@@ -27,7 +27,7 @@ from .compat import add_op_namespace_prefix, FP8_DTYPE, MX_SCALE_GROUP_K, NIBBLE
27
  from .descriptors import rebind_batched_mx_bs_descriptor
28
  from .formats import check_activation_format, global_scale_stride, normalize_global_scale, e2m1_as_uint8, expert_weight_shape, is_mx, mx_scale_family, normalize_per_expert_scale, resolve_activation_format, resolve_output_dtype, ue8m0_as_uint8, validate_dense_operands, weight_block_size, weight_format
29
  from .epilogue import fused_glu
30
- from .quant import MX_ACT_QUANT, fp8_act_quant_block_dynamic, fp8_act_quant_tensor_wide
31
  from .swizzle import swizzled_scale_descriptor
32
  from .mma import block_dynamic_dot, fp8_dot, mx_compute, mx_weight_only_compute, static_dot
33
  from .loading.tiles import (
@@ -44,7 +44,7 @@ from .loading.tiles import (
44
  oriented_tile_ptrs,
45
  weight_tile_ptrs,
46
  )
47
- from .epilogue import acc_finalize, acc_init, add_bias, bias_strides, gemm_epilogue
48
  from .pruners import PATH_ANCHOR_AXES, fp8_dot_warp_pruner, dot_scaled_staging_pruner, block_fits_dim_pruner, block_within_dim_pruner, compose_pruners, mx_config_pruner, require_moe_dims_aligned, scale_subblock_pruner, smem_pruner, swizzled_scale_config_pruner, weight_only_swap_scope_pruner
49
 
50
 
@@ -94,30 +94,6 @@ def expert_setup(
94
  return batch_id, pid_n, expert_id, A, B, C, Bs, in_row, out_row
95
 
96
 
97
- @triton.jit
98
- def store_row(
99
- C,
100
- accumulator,
101
- pid_n,
102
- stride_c_n,
103
- BLOCK_SIZE_M: tl.constexpr,
104
- BLOCK_SIZE_N: tl.constexpr,
105
- ):
106
- """Output epilogue shared by the batched kernels (``C`` already advanced to the
107
- row). The fake-batch trick aliases all ``BLOCK_SIZE_M`` lanes to the same C row,
108
- so a plain store would issue ``BLOCK_SIZE_M`` duplicate-address writes — benign on
109
- NVIDIA WGMMA (last-write-wins of identical bytes) but hardware-undefined on Intel
110
- XPU, where it corrupts the output. Mask so only lane 0 stores; the accumulator
111
- rows are mathematically identical (same A row × same B), so lane 0 is correct."""
112
- c = accumulator.to(C.dtype.element_ty)
113
- offs_cm = tl.arange(0, BLOCK_SIZE_M)
114
- offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
115
- # offs_cm[:, None] * 0: broadcast to a [BM, BN] pointer tile (all rows alias the one C row)
116
- # so the lane-0 mask below has a row axis to select; the M stride is deliberately 0.
117
- c_ptrs = C + offs_cm[:, None] * 0 + stride_c_n * offs_cn[None, :]
118
- tl.store(c_ptrs, c, mask=(offs_cm == 0)[:, None])
119
-
120
-
121
  @bayesian_autotune(
122
  get_accelerator_autotuning_configs(swap_ab=True, tune_block_n=True),
123
  # one winner per (shape, requant): requant narrows the legal tiles to the quant block, so a
@@ -281,7 +257,7 @@ def w8a8_block_dynamic_fp8_matmul_batched_kernel(
281
  @triton.jit
282
  def w8a8_block_static_fp8_matmul_batched_kernel(
283
  A, # (S, K) E4M3 activations (pre-quantized against the static scale by the wrapper)
284
- As, # scalar — static per-tensor activation scale (calibration-time)
285
  B, # (num_experts, N, K) FP8 weights; under GATE the (num_experts, 2N, K) gate|up stack
286
  Bs, # (num_experts, N // BLOCK_SIZE_N, K // BLOCK_SIZE_K) weight scales (2N under GATE)
287
  C, # (S, N) output; under an OUTPUT_FORMAT the FP8-requantized intermediate
@@ -300,6 +276,7 @@ def w8a8_block_static_fp8_matmul_batched_kernel(
300
  stride_b_e,
301
  stride_b_k,
302
  stride_b_n,
 
303
  stride_bs_e,
304
  stride_bs_k,
305
  stride_bs_n,
@@ -334,7 +311,6 @@ def w8a8_block_static_fp8_matmul_batched_kernel(
334
  pre-quantized against the calibrated scalar, per-block weight scales apply per-K-tile
335
  (``accumulate`` ``"static"``, ``FAKE_BATCH``), and the scalar activation scale multiplies the
336
  accumulator once after the loop. bf16 GLU output only (no fused requant). GATE=False is the plain GEMM."""
337
- a_s_static = tl.load(As) # per-tensor static activation scale, applied post-loop
338
  if PDL:
339
  gdc_wait()
340
  batch_id, pid_n, expert_id, A, B, C, Bs, in_row, out_row = expert_setup(
@@ -355,6 +331,10 @@ def w8a8_block_static_fp8_matmul_batched_kernel(
355
  if expert_id >= num_experts:
356
  return
357
 
 
 
 
 
358
  # the N tile may subdivide the quant block — see the dynamic sibling / scale_subblock_pruner
359
  n_width: tl.constexpr = 2 * BLOCK_SIZE_N if GATE else BLOCK_SIZE_N
360
  offs_bn = pid_n * n_width + tl.arange(0, n_width)
@@ -368,7 +348,9 @@ def w8a8_block_static_fp8_matmul_batched_kernel(
368
  acc = acc_init("dot", BLOCK_SIZE_M, n_width, SWAP_AB)
369
 
370
  for k in range(0, tl.cdiv(K, BLOCK_SIZE_K)):
371
- a, _ = load_act_static(a_ptrs, 0, 0, 0, None, 0, 0.0, "pointer", False) # pre-quantized E4M3 token (fake-batch replicated)
 
 
372
  w, b_s = load_weight_static(
373
  b_ptrs, b_ptrs, bs_ptr + bs_off, None, 0, 0, 0, 0, 0, 0,
374
  GATE, False, "pointer", SWAP_AB, BLOCK_SIZE_N, BLOCK_SIZE_K,
@@ -412,10 +394,10 @@ def w8a8_block_static_fp8_matmul_batched_kernel(
412
  @triton.jit
413
  def w8a8_tensor_dynamic_fp8_matmul_batched_kernel(
414
  A, # (S, K) pre-quantized FP8 activations
415
- As, # (S,) per-token activation scales
416
- B, # (num_experts, N, K) FP8 weight matrices
417
  Bs, # (num_experts, 1, 1) per-tensor weight scales
418
- C, # (S, N) output
419
  Bias, # (E, N_out) per-expert output bias, N_out = 2N under GATE; read iff not None
420
  ExpertIds, # (S,) — which expert each batch element routes to
421
  GatherIdx, # (S,) int — batch_id -> source row of A; read only when not None
@@ -428,6 +410,7 @@ def w8a8_tensor_dynamic_fp8_matmul_batched_kernel(
428
  stride_a_m,
429
  stride_a_k,
430
  stride_as_m,
 
431
  stride_b_e,
432
  stride_b_k,
433
  stride_b_n,
@@ -443,6 +426,13 @@ def w8a8_tensor_dynamic_fp8_matmul_batched_kernel(
443
  BLOCK_SIZE_N: tl.constexpr,
444
  BLOCK_SIZE_K: tl.constexpr,
445
  SWAP_AB: tl.constexpr = False,
 
 
 
 
 
 
 
446
  PDL: tl.constexpr = False,
447
  ):
448
  """Tensor-scale batched FP8 expert matmul kernel.
@@ -452,7 +442,10 @@ def w8a8_tensor_dynamic_fp8_matmul_batched_kernel(
452
 
453
  ``SWAP_AB`` (tuner axis, M=1 decode): weight output rows in the MMA M dim (``B`` as ``[BN, BK]``,
454
  single token padded to N=16); column 0 of the ``[BN, 16]`` accumulator is the result. Both
455
- scales are per-token/per-tensor scalars, applied once after the loop, orientation-agnostic."""
 
 
 
456
  if PDL:
457
  gdc_wait()
458
  batch_id, pid_n, expert_id, A, B, C, Bs, in_row, out_row = expert_setup(
@@ -473,18 +466,20 @@ def w8a8_tensor_dynamic_fp8_matmul_batched_kernel(
473
  if expert_id >= num_experts:
474
  return
475
 
476
- offs_bn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
 
 
477
  offs_k = tl.arange(0, BLOCK_SIZE_K)
478
  a_ptrs = operand_tile_ptrs(A, tl.arange(0, BLOCK_SIZE_M) * 0, offs_k, stride_a_m, stride_a_k, "pointer", True)
479
  b_ptrs = oriented_tile_ptrs(B, offs_bn, offs_k, stride_b_n, stride_b_k, SWAP_AB)
480
  b_s = tl.load(Bs)
481
- a_s = tl.load(As + in_row * stride_as_m)
482
 
483
- accumulator = acc_init("dot", BLOCK_SIZE_M, BLOCK_SIZE_N, SWAP_AB)
484
  for _ in range(0, tl.cdiv(K, BLOCK_SIZE_K)):
485
- a, _ = load_act_plain(a_ptrs, 0, 0, 0, None, 0, "pointer", False)
486
  b, _ = load_weight_plain(
487
- b_ptrs, b_ptrs, 0, 0, 0, False, False, "pointer", SWAP_AB, BLOCK_SIZE_N, BLOCK_SIZE_K
488
  )
489
  accumulator = accumulator + fp8_dot(a, b, SWAP_AB, BLOCK_SIZE_K)
490
  a_ptrs, _, b_ptrs, _, _, _ = advance_ptrs(
@@ -493,14 +488,21 @@ def w8a8_tensor_dynamic_fp8_matmul_batched_kernel(
493
  "pointer", "pointer", False, False, False,
494
  )
495
 
496
- accumulator = acc_finalize(accumulator, "dot", BLOCK_SIZE_N, SWAP_AB) * a_s * b_s
497
- # this split keeps its own store path (ungated, per-tensor dequant) but takes the same bias
498
- accumulator = add_bias(
499
- accumulator, Bias, stride_bias_e, stride_bias_n, expert_id, pid_n, BLOCK_SIZE_N
500
- )
501
  if PDL:
502
  gdc_launch_dependents()
503
- store_row(C, accumulator, pid_n, stride_c_n, BLOCK_SIZE_M, BLOCK_SIZE_N)
 
 
 
 
 
 
 
504
 
505
 
506
  # The MXFP4/MXFP8 (and packed-activation) splits key themselves — the tuner appends every tensor
@@ -1170,7 +1172,8 @@ def w8a8_block_static_fp8_matmul_batched(
1170
 
1171
  A: (rows, K) raw bf16/fp16 activations — rows addressed via ``gather_idx``
1172
  B: (num_experts, N, K) FP8 weights; under ``gate`` the (num_experts, 2N, K) gate|up stack
1173
- As: scalar / (1,) — the calibrated per-tensor (static) activation scale
 
1174
  Bs: (num_experts, N // block_n, K // block_k) per-block weight scales (2N under gate)
1175
  """
1176
  validate_dense_operands(A, B)
@@ -1197,11 +1200,8 @@ def w8a8_block_static_fp8_matmul_batched(
1197
  f"the fused 'fp8' requant needs square quant blocks, got {block_size}"
1198
  )
1199
 
1200
- As = As.reshape(1).to(torch.float32)
1201
  bs_u8 = ue8m0_as_uint8(Bs)
1202
- # Pre-quantize the raw activations against the calibrated scalar (offline; the kernel folds
1203
- # the scalar back post-loop).
1204
- A_q = (A.to(torch.float32) / As).to(FP8_DTYPE)
1205
  if requant:
1206
  C = A.new_empty(S, N, dtype=FP8_DTYPE)
1207
  Cs = torch.empty(S, N // block_n, device=A.device, dtype=bs_u8.dtype)
@@ -1235,6 +1235,7 @@ def w8a8_block_static_fp8_matmul_batched(
1235
  B.stride(0),
1236
  B.stride(2),
1237
  B.stride(1),
 
1238
  bs_u8.stride(0),
1239
  bs_u8.stride(2),
1240
  bs_u8.stride(1),
@@ -1274,6 +1275,11 @@ def w8a8_tensor_dynamic_fp8_matmul_batched(
1274
  As: torch.Tensor | None,
1275
  Bs: torch.Tensor,
1276
  expert_ids: torch.Tensor,
 
 
 
 
 
1277
  output_dtype: torch.dtype | None = None,
1278
  gather_idx: torch.Tensor | None = None,
1279
  scatter_idx: torch.Tensor | None = None,
@@ -1285,8 +1291,9 @@ def w8a8_tensor_dynamic_fp8_matmul_batched(
1285
  (None = row s).
1286
 
1287
  A: (rows, K) raw or pre-quantized FP8 activations — rows addressed via ``gather_idx``
1288
- B: (num_experts, N, K) FP8 expert weights
1289
- As: (rows,) per-token scales, or None when A is raw
 
1290
  Bs: (num_experts,) or (num_experts, 1, 1) per-expert weight scales
1291
  """
1292
  validate_dense_operands(A, B)
@@ -1294,16 +1301,15 @@ def w8a8_tensor_dynamic_fp8_matmul_batched(
1294
  output_dtype = resolve_output_dtype(output_dtype, A, As)
1295
  K = A.shape[1]
1296
  S = expert_ids.shape[0]
1297
- num_experts, N, _ = B.shape
 
 
1298
 
1299
  # Normalize Bs to (num_experts, 1, 1)
1300
  Bs = normalize_per_expert_scale(Bs, num_experts)
1301
 
1302
  bs_u8 = ue8m0_as_uint8(Bs)
1303
- if As is None:
1304
- qA, As = fp8_act_quant_tensor_wide(A, K)
1305
- else:
1306
- qA = A
1307
  C = A.new_empty(S, N, dtype=output_dtype)
1308
 
1309
  def grid(META):
@@ -1328,7 +1334,8 @@ def w8a8_tensor_dynamic_fp8_matmul_batched(
1328
  K,
1329
  qA.stride(0),
1330
  qA.stride(1),
1331
- As.stride(0),
 
1332
  B.stride(0),
1333
  B.stride(2),
1334
  B.stride(1),
@@ -1339,6 +1346,11 @@ def w8a8_tensor_dynamic_fp8_matmul_batched(
1339
  bias_stride_n,
1340
  expert_ids.stride(0),
1341
  num_experts=num_experts,
 
 
 
 
 
1342
  PDL=decode_pdl(),
1343
  launch_pdl=decode_pdl(),
1344
  )
@@ -1819,12 +1831,15 @@ def matmul_batched(
1819
  "output_global_scale is the NVFP4 requant second level — it requires quantize_output=True on NVFP4 "
1820
  "(the epilogue would otherwise normalize by it with nothing downstream to compensate)"
1821
  )
1822
- if As is not None and As.numel() == 1:
1823
- # static (per-tensor calibrated) activation quant: a per-tensor scalar As for block-scale FP8
1824
- # weights — the caller hands raw A, the op quantizes it against the scalar (As IS the scale).
1825
- assert Bs is not None and not is_mx(B, Bs) and weight_block_size(B, Bs) is not None, (
1826
- "a per-tensor scalar As (static activation scale) needs block-scale FP8 weights"
 
 
1827
  )
 
1828
  out = w8a8_block_static_fp8_matmul_batched(
1829
  A,
1830
  B,
@@ -1912,14 +1927,12 @@ def matmul_batched(
1912
  bias=bias,
1913
  )
1914
  elif (block_size := weight_block_size(B, Bs)) is None:
1915
- assert not gate, (
1916
- "the batched op has no tensor-wide gate|up fusion (grouped and 2D support it)"
1917
- )
1918
  assert activation_format in (None, "fp8") and not quantize_output, (
1919
  "tensor-wide supports neither packed activations nor a fused requant"
1920
  )
1921
  out = w8a8_tensor_dynamic_fp8_matmul_batched(
1922
- A, B, As, Bs, expert_ids, output_dtype, gather_idx, scatter_idx, bias=bias
 
1923
  )
1924
  else:
1925
  out = w8a8_block_dynamic_fp8_matmul_batched(
 
27
  from .descriptors import rebind_batched_mx_bs_descriptor
28
  from .formats import check_activation_format, global_scale_stride, normalize_global_scale, e2m1_as_uint8, expert_weight_shape, is_mx, mx_scale_family, normalize_per_expert_scale, resolve_activation_format, resolve_output_dtype, ue8m0_as_uint8, validate_dense_operands, weight_block_size, weight_format
29
  from .epilogue import fused_glu
30
+ from .quant import fp8_act_quant_block_dynamic, MX_ACT_QUANT, static_expert_act_operands, tensor_wide_act_operands
31
  from .swizzle import swizzled_scale_descriptor
32
  from .mma import block_dynamic_dot, fp8_dot, mx_compute, mx_weight_only_compute, static_dot
33
  from .loading.tiles import (
 
44
  oriented_tile_ptrs,
45
  weight_tile_ptrs,
46
  )
47
+ from .epilogue import acc_init, bias_strides, gemm_epilogue
48
  from .pruners import PATH_ANCHOR_AXES, fp8_dot_warp_pruner, dot_scaled_staging_pruner, block_fits_dim_pruner, block_within_dim_pruner, compose_pruners, mx_config_pruner, require_moe_dims_aligned, scale_subblock_pruner, smem_pruner, swizzled_scale_config_pruner, weight_only_swap_scope_pruner
49
 
50
 
 
94
  return batch_id, pid_n, expert_id, A, B, C, Bs, in_row, out_row
95
 
96
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
97
  @bayesian_autotune(
98
  get_accelerator_autotuning_configs(swap_ab=True, tune_block_n=True),
99
  # one winner per (shape, requant): requant narrows the legal tiles to the quant block, so a
 
257
  @triton.jit
258
  def w8a8_block_static_fp8_matmul_batched_kernel(
259
  A, # (S, K) E4M3 activations (pre-quantized against the static scale by the wrapper)
260
+ As, # calibrated (static) activation scale: one value, or one per expert
261
  B, # (num_experts, N, K) FP8 weights; under GATE the (num_experts, 2N, K) gate|up stack
262
  Bs, # (num_experts, N // BLOCK_SIZE_N, K // BLOCK_SIZE_K) weight scales (2N under GATE)
263
  C, # (S, N) output; under an OUTPUT_FORMAT the FP8-requantized intermediate
 
276
  stride_b_e,
277
  stride_b_k,
278
  stride_b_n,
279
+ stride_as_e, # 0 = one calibrated scale shared by every expert
280
  stride_bs_e,
281
  stride_bs_k,
282
  stride_bs_n,
 
311
  pre-quantized against the calibrated scalar, per-block weight scales apply per-K-tile
312
  (``accumulate`` ``"static"``, ``FAKE_BATCH``), and the scalar activation scale multiplies the
313
  accumulator once after the loop. bf16 GLU output only (no fused requant). GATE=False is the plain GEMM."""
 
314
  if PDL:
315
  gdc_wait()
316
  batch_id, pid_n, expert_id, A, B, C, Bs, in_row, out_row = expert_setup(
 
331
  if expert_id >= num_experts:
332
  return
333
 
334
+ # this batch's expert scale: stride 0 means one calibrated scale for every expert. It feeds
335
+ # the inline quant arm and folds back onto the accumulator post-loop.
336
+ a_s_static = tl.load(As + expert_id.to(tl.int32) * stride_as_e)
337
+
338
  # the N tile may subdivide the quant block — see the dynamic sibling / scale_subblock_pruner
339
  n_width: tl.constexpr = 2 * BLOCK_SIZE_N if GATE else BLOCK_SIZE_N
340
  offs_bn = pid_n * n_width + tl.arange(0, n_width)
 
348
  acc = acc_init("dot", BLOCK_SIZE_M, n_width, SWAP_AB)
349
 
350
  for k in range(0, tl.cdiv(K, BLOCK_SIZE_K)):
351
+ # raw A quantizes in register against this expert's scale; a pre-quantized E4M3 token
352
+ # (fake-batch replicated) takes the other arm and ignores it
353
+ a, _ = load_act_static(a_ptrs, 0, 0, 0, None, 0, a_s_static, "pointer", False)
354
  w, b_s = load_weight_static(
355
  b_ptrs, b_ptrs, bs_ptr + bs_off, None, 0, 0, 0, 0, 0, 0,
356
  GATE, False, "pointer", SWAP_AB, BLOCK_SIZE_N, BLOCK_SIZE_K,
 
394
  @triton.jit
395
  def w8a8_tensor_dynamic_fp8_matmul_batched_kernel(
396
  A, # (S, K) pre-quantized FP8 activations
397
+ As, # per-token activation scales (S,), or a calibrated (static) one: shared, or per expert
398
+ B, # (num_experts, N, K) FP8 weights; under GATE the (num_experts, 2N, K) gate|up stack
399
  Bs, # (num_experts, 1, 1) per-tensor weight scales
400
+ C, # (S, N) output; under GATE the bf16 GLU intermediate
401
  Bias, # (E, N_out) per-expert output bias, N_out = 2N under GATE; read iff not None
402
  ExpertIds, # (S,) — which expert each batch element routes to
403
  GatherIdx, # (S,) int — batch_id -> source row of A; read only when not None
 
410
  stride_a_m,
411
  stride_a_k,
412
  stride_as_m,
413
+ stride_as_e, # 0 = the scale is per token; 1 = one calibrated scale per expert
414
  stride_b_e,
415
  stride_b_k,
416
  stride_b_n,
 
426
  BLOCK_SIZE_N: tl.constexpr,
427
  BLOCK_SIZE_K: tl.constexpr,
428
  SWAP_AB: tl.constexpr = False,
429
+ # Gate|up fusion epilogue (GATE=False -> plain batched GEMM, every arm below folds out)
430
+ GATE: tl.constexpr = False,
431
+ ACT_FN: tl.constexpr = "silu",
432
+ SWIGLU_ALPHA: tl.constexpr = None,
433
+ SWIGLU_LIMIT: tl.constexpr = None,
434
+ SIMULATE_UNFUSED: tl.constexpr = False,
435
+ INTERMEDIATE_DTYPE: tl.constexpr = tl.bfloat16,
436
  PDL: tl.constexpr = False,
437
  ):
438
  """Tensor-scale batched FP8 expert matmul kernel.
 
442
 
443
  ``SWAP_AB`` (tuner axis, M=1 decode): weight output rows in the MMA M dim (``B`` as ``[BN, BK]``,
444
  single token padded to N=16); column 0 of the ``[BN, 16]`` accumulator is the result. Both
445
+ scales are per-token/per-tensor scalars, applied once after the loop, orientation-agnostic.
446
+ ``GATE`` fuses the gate|up projection (``B`` the interleaved (2N, K) stack, one per-tensor
447
+ scale over both halves) into one tile + SwiGLU, emitting the bf16 intermediate; ``GATE=False``
448
+ is the plain GEMM (bit-identical)."""
449
  if PDL:
450
  gdc_wait()
451
  batch_id, pid_n, expert_id, A, B, C, Bs, in_row, out_row = expert_setup(
 
466
  if expert_id >= num_experts:
467
  return
468
 
469
+ # under GATE the gate|up rows are interleaved, so the tile is one 2*BN span
470
+ n_width: tl.constexpr = 2 * BLOCK_SIZE_N if GATE else BLOCK_SIZE_N
471
+ offs_bn = pid_n * n_width + tl.arange(0, n_width)
472
  offs_k = tl.arange(0, BLOCK_SIZE_K)
473
  a_ptrs = operand_tile_ptrs(A, tl.arange(0, BLOCK_SIZE_M) * 0, offs_k, stride_a_m, stride_a_k, "pointer", True)
474
  b_ptrs = oriented_tile_ptrs(B, offs_bn, offs_k, stride_b_n, stride_b_k, SWAP_AB)
475
  b_s = tl.load(Bs)
476
+ a_s = tl.load(As + in_row * stride_as_m + expert_id * stride_as_e)
477
 
478
+ accumulator = acc_init("dot", BLOCK_SIZE_M, n_width, SWAP_AB)
479
  for _ in range(0, tl.cdiv(K, BLOCK_SIZE_K)):
480
+ a, _ = load_act_static(a_ptrs, 0, 0, 0, None, 0, a_s, "pointer", False)
481
  b, _ = load_weight_plain(
482
+ b_ptrs, b_ptrs, 0, 0, 0, GATE, False, "pointer", SWAP_AB, BLOCK_SIZE_N, BLOCK_SIZE_K
483
  )
484
  accumulator = accumulator + fp8_dot(a, b, SWAP_AB, BLOCK_SIZE_K)
485
  a_ptrs, _, b_ptrs, _, _, _ = advance_ptrs(
 
488
  "pointer", "pointer", False, False, False,
489
  )
490
 
491
+ # per-tensor scales fold on the raw accumulator; `gemm_epilogue` owns the finalize.
492
+ # Finalizing here too was invisible under no-swap (a pass-through) and collapsed the
493
+ # swapped tile twice, so every SWAP_AB config failed to compile and the tuner quietly
494
+ # dropped the whole swap arm at this shape.
495
+ accumulator = accumulator * a_s * b_s
496
  if PDL:
497
  gdc_launch_dependents()
498
+ gemm_epilogue(
499
+ C, None, accumulator, out_row, pid_n, 0, out_row, 1, stride_c_n, 1, 1,
500
+ BLOCK_SIZE_M, BLOCK_SIZE_N, GATE, None, 1,
501
+ ACT_FN, SWIGLU_ALPHA, SWIGLU_LIMIT, SIMULATE_UNFUSED, INTERMEDIATE_DTYPE,
502
+ COMPUTE_MODE="dot", SWAP_AB=SWAP_AB, FAKE_BATCH=True,
503
+ Bias=Bias, stride_bias_e=stride_bias_e, stride_bias_n=stride_bias_n,
504
+ global_row=expert_id,
505
+ )
506
 
507
 
508
  # The MXFP4/MXFP8 (and packed-activation) splits key themselves — the tuner appends every tensor
 
1172
 
1173
  A: (rows, K) raw bf16/fp16 activations — rows addressed via ``gather_idx``
1174
  B: (num_experts, N, K) FP8 weights; under ``gate`` the (num_experts, 2N, K) gate|up stack
1175
+ As: scalar / (1,) for one calibrated scale, or (num_experts,) for a MoE that
1176
+ calibrates each expert separately
1177
  Bs: (num_experts, N // block_n, K // block_k) per-block weight scales (2N under gate)
1178
  """
1179
  validate_dense_operands(A, B)
 
1200
  f"the fused 'fp8' requant needs square quant blocks, got {block_size}"
1201
  )
1202
 
1203
+ A_q, As, as_stride = static_expert_act_operands(A, As, num_experts)
1204
  bs_u8 = ue8m0_as_uint8(Bs)
 
 
 
1205
  if requant:
1206
  C = A.new_empty(S, N, dtype=FP8_DTYPE)
1207
  Cs = torch.empty(S, N // block_n, device=A.device, dtype=bs_u8.dtype)
 
1235
  B.stride(0),
1236
  B.stride(2),
1237
  B.stride(1),
1238
+ as_stride,
1239
  bs_u8.stride(0),
1240
  bs_u8.stride(2),
1241
  bs_u8.stride(1),
 
1275
  As: torch.Tensor | None,
1276
  Bs: torch.Tensor,
1277
  expert_ids: torch.Tensor,
1278
+ gate: bool = False,
1279
+ act_fn: str = "silu",
1280
+ swiglu_alpha: float | None = None,
1281
+ swiglu_limit: float | None = None,
1282
+ simulate_unfused: bool = False,
1283
  output_dtype: torch.dtype | None = None,
1284
  gather_idx: torch.Tensor | None = None,
1285
  scatter_idx: torch.Tensor | None = None,
 
1291
  (None = row s).
1292
 
1293
  A: (rows, K) raw or pre-quantized FP8 activations — rows addressed via ``gather_idx``
1294
+ B: (num_experts, N, K) FP8 expert weights; under ``gate`` the (num_experts, 2N, K) stack
1295
+ As: (rows,) per-token scales alongside a pre-quantized A, or — on a raw A — the calibrated
1296
+ (static) activation scale: one value, or one per expert
1297
  Bs: (num_experts,) or (num_experts, 1, 1) per-expert weight scales
1298
  """
1299
  validate_dense_operands(A, B)
 
1301
  output_dtype = resolve_output_dtype(output_dtype, A, As)
1302
  K = A.shape[1]
1303
  S = expert_ids.shape[0]
1304
+ num_experts, rows, _ = B.shape
1305
+ # under gate|up fusion B is the (E, 2N, K) stack; N is the per-projection output width
1306
+ N = rows // 2 if gate else rows
1307
 
1308
  # Normalize Bs to (num_experts, 1, 1)
1309
  Bs = normalize_per_expert_scale(Bs, num_experts)
1310
 
1311
  bs_u8 = ue8m0_as_uint8(Bs)
1312
+ qA, As, as_stride_m, as_stride_e = tensor_wide_act_operands(A, As, num_experts)
 
 
 
1313
  C = A.new_empty(S, N, dtype=output_dtype)
1314
 
1315
  def grid(META):
 
1334
  K,
1335
  qA.stride(0),
1336
  qA.stride(1),
1337
+ as_stride_m,
1338
+ as_stride_e,
1339
  B.stride(0),
1340
  B.stride(2),
1341
  B.stride(1),
 
1346
  bias_stride_n,
1347
  expert_ids.stride(0),
1348
  num_experts=num_experts,
1349
+ GATE=gate,
1350
+ ACT_FN=act_fn,
1351
+ SWIGLU_ALPHA=swiglu_alpha,
1352
+ SWIGLU_LIMIT=swiglu_limit,
1353
+ SIMULATE_UNFUSED=simulate_unfused,
1354
  PDL=decode_pdl(),
1355
  launch_pdl=decode_pdl(),
1356
  )
 
1831
  "output_global_scale is the NVFP4 requant second level — it requires quantize_output=True on NVFP4 "
1832
  "(the epilogue would otherwise normalize by it with nothing downstream to compensate)"
1833
  )
1834
+ # a calibrated (static) scale on a raw A (see `tensor_wide_act_operands`): block-scale weights
1835
+ # have a dedicated static kernel, per-tensor ones read the same As on the tensor-wide arm
1836
+ static_act = As is not None and As.ndim <= 1 and As.numel() in (1, B.shape[0]) and A.dtype != FP8_DTYPE
1837
+ if static_act:
1838
+ assert Bs is not None and not is_mx(B, Bs), (
1839
+ "a calibrated (static) activation scale is an FP8 form — MX activations carry a scale "
1840
+ "per group, derived per call"
1841
  )
1842
+ if static_act and weight_block_size(B, Bs) is not None:
1843
  out = w8a8_block_static_fp8_matmul_batched(
1844
  A,
1845
  B,
 
1927
  bias=bias,
1928
  )
1929
  elif (block_size := weight_block_size(B, Bs)) is None:
 
 
 
1930
  assert activation_format in (None, "fp8") and not quantize_output, (
1931
  "tensor-wide supports neither packed activations nor a fused requant"
1932
  )
1933
  out = w8a8_tensor_dynamic_fp8_matmul_batched(
1934
+ A, B, As, Bs, expert_ids, gate, act_fn, swiglu_alpha, swiglu_limit, simulate_unfused,
1935
+ output_dtype, gather_idx, scatter_idx, bias=bias,
1936
  )
1937
  else:
1938
  out = w8a8_block_dynamic_fp8_matmul_batched(
build/torch-rocm/bayesian_autotuner.py CHANGED
@@ -101,11 +101,12 @@ class BayesianAutotuner(Autotuner):
101
  # budget being met. Defaults to n_trials: a grid that cannot land its measurements without
102
  # that many rejects has a pruner gap, and the fix is to fence the dead region rather than
103
  # to tolerate the compiles. Its own axis so a kernel can raise it deliberately.
104
- self.max_failures = int(
105
- os.environ.get("FINEGRAINED_AUTOTUNE_MAX_FAILURES")
106
- or max_failures
107
- or self.n_trials
108
- )
 
109
  self.n_startup_trials = n_startup_trials
110
  # top fraction of measured configs the TPE treats as "good"
111
  self.gamma = gamma
 
101
  # budget being met. Defaults to n_trials: a grid that cannot land its measurements without
102
  # that many rejects has a pruner gap, and the fix is to fence the dead region rather than
103
  # to tolerate the compiles. Its own axis so a kernel can raise it deliberately.
104
+ # an explicit argument wins over the env override, which wins over the default; `or`
105
+ # would drop a deliberate `max_failures=0`, the strictest setting there is
106
+ if max_failures is None:
107
+ env_max_failures = os.environ.get("FINEGRAINED_AUTOTUNE_MAX_FAILURES")
108
+ max_failures = env_max_failures if env_max_failures is not None else self.n_trials
109
+ self.max_failures = int(max_failures)
110
  self.n_startup_trials = n_startup_trials
111
  # top fraction of measured configs the TPE treats as "good"
112
  self.gamma = gamma
build/torch-rocm/compat.py CHANGED
@@ -40,6 +40,14 @@ def decode_pdl() -> bool:
40
  no fusion). Deployment decode runs eager launches under cudagraphs, where PDL applies."""
41
  return DECODE_PDL and not torch.compiler.is_compiling()
42
 
 
 
 
 
 
 
 
 
43
  # kernel-builder generates ``_ops.py`` into the built variant (the op namespace carries the build's
44
  # unique id) and refuses to build over a tracked one, so the source tree has none: an unbuilt
45
  # checkout (tests, bench, FINEGRAINED_KERNELS_PATH) registers its ops under the plain package name.
 
40
  no fusion). Deployment decode runs eager launches under cudagraphs, where PDL applies."""
41
  return DECODE_PDL and not torch.compiler.is_compiling()
42
 
43
+ # The scaled_grouped_mm scaling enums the torch MoE baseline hands its scale operands. Older
44
+ # torch has neither them nor the op, and only that one baseline reads them, so a missing pair is
45
+ # ``None`` here and the baseline refuses on it rather than the whole package failing to import.
46
+ try:
47
+ from torch.nn.functional import ScalingType, SwizzleType
48
+ except ImportError:
49
+ ScalingType = SwizzleType = None
50
+
51
  # kernel-builder generates ``_ops.py`` into the built variant (the op namespace carries the build's
52
  # unique id) and refuses to build over a tracked one, so the source tree has none: an unbuilt
53
  # checkout (tests, bench, FINEGRAINED_KERNELS_PATH) registers its ops under the plain package name.
build/torch-rocm/grouped.py CHANGED
@@ -24,10 +24,10 @@ from .bayesian_autotuner import bayesian_autotune
24
  from .compat import add_op_namespace_prefix, FP8_DTYPE, NIBBLES_PER_BYTE, compile_time_only_triton_op, compile_time_only_triton_wrap, device_context, get_accelerator_autotuning_configs, sm_count, tl_dtype
25
  from .descriptors import build_grouped_operand_descriptors, rebind_grouped_descriptors, rebind_grouped_mx_descriptors
26
  from .formats import check_activation_format, global_scale_stride, is_per_expert_global, normalize_global_scale, e2m1_as_uint8, expert_weight_shape, is_mx, mx_scale_family, normalize_per_expert_scale, resolve_activation_format, resolve_output_dtype, routed_rows, tokens_per_expert_bucket, ue8m0_as_uint8, validate_dense_operands, weight_block_size, weight_format
27
- from .quant import MX_ACT_QUANT, fp8_act_quant_block_dynamic, fp8_act_quant_tensor_wide, mx_act_quant_grouped, swizzle_grouped_mx_scales
28
  from .swizzle import swizzled_scale_descriptor
29
  from .mma import block_dynamic_dot, fp8_dot, mx_compute, mx_weight_only_compute, static_dot
30
- from .scheduling import build_tile_layout, expand_gather_below_parity, build_packed_schedule, load_packed_schedule, prefetch_packed_entry, resolve_grouped_tile, resolve_grouped_tile_packed
31
  from .loading.tiles import (
32
  load_act_block_dynamic,
33
  load_act_mx,
@@ -41,7 +41,7 @@ from .loading.tiles import (
41
  weight_tile_ptrs,
42
  )
43
  from .epilogue import acc_init, bias_strides, gemm_epilogue
44
- from .pruners import PATH_ANCHOR_AXES, fp8_dot_warp_pruner, global_scale_warp_spec_pruner, weight_only_warp_spec_matched_mode_pruner, packed_schedule_scope_pruner, affine_scale_warp_spec_pruner, block_dynamic_grouped_matmul_pruner, block_fits_dim_pruner, block_within_dim_pruner, compose_pruners, descriptor_box_pruner, gate_stacked_tmem_trap_pruner, gated_pointer_weight_warp_spec_pruner, mx_config_pruner, require_moe_dims_aligned, smem_pruner, swizzled_out_bm_pruner, swizzled_scale_config_pruner, swizzled_scales_bm_pruner, warp_spec_compile_guard_pruner
45
 
46
 
47
  @bayesian_autotune(
@@ -69,6 +69,7 @@ from .pruners import PATH_ANCHOR_AXES, fp8_dot_warp_pruner, global_scale_warp_sp
69
  fp8_dot_warp_pruner(),
70
  packed_schedule_scope_pruner(),
71
  block_dynamic_grouped_matmul_pruner(),
 
72
  descriptor_box_pruner(),
73
  )
74
  },
@@ -274,6 +275,7 @@ def w8a8_block_dynamic_fp8_matmul_grouped_kernel(
274
  packed_schedule_scope_pruner(),
275
  block_dynamic_grouped_matmul_pruner(),
276
  descriptor_box_pruner(),
 
277
  )
278
  },
279
  )
@@ -281,7 +283,7 @@ def w8a8_block_dynamic_fp8_matmul_grouped_kernel(
281
  def w8a8_block_static_fp8_matmul_grouped_kernel(
282
  A, # (num_tokens, K) E4M3 activations (pre-quantized against the static scale by the wrapper)
283
  ADescriptor, # host TMA descriptor over A (rows, K), box (BM, BK); read iff A_MEMORY_MODE != "pointer"
284
- As, # scalar — static per-tensor activation scale (calibration-time)
285
  B, # (num_experts, N, K) FP8 weights; under GATE the (num_experts, 2N, K) gate|up stack
286
  BDescriptor, # host TMA descriptor over B viewed (E, 2N|N, K), box (1, (2|1)*BN, BK); read iff B_MEMORY_MODE != "pointer"
287
  Bs, # (num_experts, N // BLOCK_SIZE_N, K // BLOCK_SIZE_K) weight scales (2N under GATE)
@@ -302,6 +304,7 @@ def w8a8_block_static_fp8_matmul_grouped_kernel(
302
  stride_b_e,
303
  stride_b_k,
304
  stride_b_n,
 
305
  stride_bs_e,
306
  stride_bs_k,
307
  stride_bs_n,
@@ -342,7 +345,6 @@ def w8a8_block_static_fp8_matmul_grouped_kernel(
342
  weight scales apply per-K-tile (plain ``tl.dot`` + software rescale, ``accumulate("static")``),
343
  and the scalar activation scale multiplies the accumulator once after the loop. GATE=False is
344
  the plain grouped GEMM (down projection), bit-identical."""
345
- a_s_static = tl.load(As) # per-tensor static activation scale, applied post-loop
346
  start_pid = tl.program_id(axis=0)
347
  if PACKED_SCHEDULE:
348
  total_m_tiles = load_packed_schedule(Schedule, BLOCK_SIZE_M)
@@ -391,10 +393,13 @@ def w8a8_block_static_fp8_matmul_grouped_kernel(
391
  gate_s_ptr = Bs + expert_id64 * stride_bs_e + ((2 if GATE else 1) * pid_n) * stride_bs_n
392
  up_s_ptr = gate_s_ptr + stride_bs_n
393
 
 
 
 
394
  acc = acc_init("dot", BLOCK_SIZE_M, (2 if GATE else 1) * BLOCK_SIZE_N, False)
395
  for k in tl.range(0, tl.cdiv(K, BLOCK_SIZE_K), warp_specialize=WARP_SPEC):
396
  a, a_dead = load_act_static(
397
- a_ptrs, ADescriptor, m_start, k * BLOCK_SIZE_K, row_mask, in_row, 0.0,
398
  A_MEMORY_MODE, GatherIdx is not None,
399
  )
400
  w, w_s = load_weight_static(
@@ -466,6 +471,7 @@ def w8a8_block_static_fp8_matmul_grouped_kernel(
466
  warp_spec_compile_guard_pruner(),
467
  descriptor_box_pruner(),
468
  smem_pruner(),
 
469
  )
470
  },
471
  )
@@ -473,7 +479,7 @@ def w8a8_block_static_fp8_matmul_grouped_kernel(
473
  def w8a8_tensor_dynamic_fp8_matmul_grouped_kernel(
474
  A, # (num_tokens, K) pre-quantized FP8 activations, any row order
475
  ADescriptor, # host TMA descriptor over A (rows, K), box (BM, BK); read iff A_MEMORY_MODE != "pointer"
476
- As, # (S,) per-token activation scales
477
  B, # (num_experts, N, K) FP8 weights; under GATE the (num_experts, 2N, K) gate|up stack
478
  BDescriptor, # host TMA descriptor over B viewed (E, 2N|N, K), box (1, (2|1)*BN, BK); read iff B_MEMORY_MODE != "pointer"
479
  Bs, # (num_experts, 1, 1) per-tensor weight scales (one scalar covers the gate|up stack)
@@ -492,6 +498,7 @@ def w8a8_tensor_dynamic_fp8_matmul_grouped_kernel(
492
  stride_a_m,
493
  stride_a_k,
494
  stride_as_m,
 
495
  stride_b_e,
496
  stride_b_k,
497
  stride_b_n,
@@ -580,13 +587,17 @@ def w8a8_tensor_dynamic_fp8_matmul_grouped_kernel(
580
  stride_b_k,
581
  False,
582
  )
583
- a_s = tl.load(As + in_row * stride_as_m, mask=row_mask, other=0.0)
 
 
 
 
584
  b_s = tl.load(Bs + expert_id64 * stride_bs_e)
585
 
586
  acc = acc_init("dot", BLOCK_SIZE_M, (2 if GATE else 1) * BLOCK_SIZE_N, False)
587
  for k in tl.range(0, tl.cdiv(K, BLOCK_SIZE_K), warp_specialize=WARP_SPEC):
588
- a, _as = load_act_plain(
589
- a_ptrs, ADescriptor, m_start, k * BLOCK_SIZE_K, row_mask, in_row,
590
  A_MEMORY_MODE, GatherIdx is not None,
591
  )
592
  w, _ws = load_weight_plain(
@@ -665,6 +676,7 @@ def w8a8_tensor_dynamic_fp8_matmul_grouped_kernel(
665
  smem_pruner(min_sets=2), # two full buffer sets: the grouped pipeline's floor (see smem_pruner)
666
  warp_spec_compile_guard_pruner(),
667
  affine_scale_warp_spec_pruner(),
 
668
  )
669
  },
670
  )
@@ -916,7 +928,7 @@ def mx_dynamic_matmul_grouped_kernel(
916
  path_anchor_axes=PATH_ANCHOR_AXES,
917
  prune_configs_by={
918
  "early_config_prune": compose_pruners(
919
- weight_only_warp_spec_matched_mode_pruner(),
920
  packed_schedule_scope_pruner(),
921
  # the tail tile slides back in bounds (see the K-loop) — BK only has to fit
922
  block_fits_dim_pruner("K"),
@@ -1436,7 +1448,8 @@ def w8a8_block_static_fp8_matmul_grouped(
1436
 
1437
  A: (S, K) raw bf16/fp16 activations — rows addressed via ``gather_idx``
1438
  B: (num_experts, N, K) FP8 weights; under ``gate`` the (num_experts, 2N, K) gate|up stack
1439
- As: scalar / (1,) — the calibrated per-tensor (static) activation scale
 
1440
  Bs: (num_experts, N // block_n, K // block_k) per-block weight scales (2N under gate)
1441
  """
1442
  validate_dense_operands(A, B)
@@ -1463,11 +1476,8 @@ def w8a8_block_static_fp8_matmul_grouped(
1463
  )
1464
 
1465
  output_dtype = resolve_output_dtype(output_dtype, A, None)
1466
- As = As.reshape(1).to(torch.float32)
1467
  bs_u8 = ue8m0_as_uint8(Bs)
1468
- # Pre-quantize the raw activations against the calibrated scalar (offline — MoE always
1469
- # pre-quants; the kernel folds the scalar back post-loop).
1470
- A_q = (A.to(torch.float32) / As).to(FP8_DTYPE)
1471
  if requant:
1472
  C = A.new_empty(S, N, dtype=FP8_DTYPE)
1473
  # UE8M0 model (ue8m0 weights) -> UE8M0 intermediate scales; the epilogue infers the format
@@ -1510,6 +1520,7 @@ def w8a8_block_static_fp8_matmul_grouped(
1510
  B.stride(0),
1511
  B.stride(2),
1512
  B.stride(1),
 
1513
  bs_u8.stride(0),
1514
  bs_u8.stride(2),
1515
  bs_u8.stride(1),
@@ -1569,7 +1580,8 @@ def w8a8_tensor_dynamic_fp8_matmul_grouped(
1569
 
1570
  A: (S, K) pre-quantized FP8 activations — rows addressed via ``gather_idx``
1571
  B: (num_experts, N, K) FP8 expert weights; under ``gate`` the (num_experts, 2N, K) stack
1572
- As: (S,) per-token activation scales
 
1573
  Bs: (num_experts,) or (num_experts, 1, 1) per-expert weight scales
1574
  expert_start: (num_experts_pow2 + 1,) int32 — cumulative sorted-row starts, S sentinel
1575
  gather_idx: optional (S,) — sorted position -> source row of A; None = A is expert-sorted
@@ -1593,11 +1605,17 @@ def w8a8_tensor_dynamic_fp8_matmul_grouped(
1593
  # Normalize Bs to (num_experts, 1, 1) — one per-tensor scale (covers the gate|up stack)
1594
  Bs = normalize_per_expert_scale(Bs, num_experts)
1595
 
1596
- # A raw (As is None) -> quantize here (offline, per-token); else pre-quantized.
1597
  output_dtype = resolve_output_dtype(output_dtype, A, As)
1598
- if As is None:
1599
- A, As = fp8_act_quant_tensor_wide(A, K)
1600
- # post-quant: trade the in-kernel gather for one packed-row copy where that wins
 
 
 
 
 
 
 
1601
  A, _, gather_idx = expand_gather_below_parity(A, None, gather_idx, num_experts)
1602
  C = A.new_empty(S, N, dtype=output_dtype)
1603
  num_sms = sm_count(A.device.index)
@@ -1631,7 +1649,8 @@ def w8a8_tensor_dynamic_fp8_matmul_grouped(
1631
  K,
1632
  A.stride(0),
1633
  A.stride(1),
1634
- As.stride(0),
 
1635
  B.stride(0),
1636
  B.stride(2),
1637
  B.stride(1),
@@ -2234,12 +2253,15 @@ def matmul_grouped(
2234
  "output_global_scale is the NVFP4 requant second level — it requires quantize_output=True on NVFP4 "
2235
  "(the epilogue would otherwise normalize by it with nothing downstream to compensate)"
2236
  )
2237
- if As is not None and As.numel() == 1:
2238
- # static (per-tensor calibrated) activation quant: a per-tensor scalar As for block-scale FP8
2239
- # weights — the caller hands raw A, the op quantizes it against the scalar (As IS the scale).
2240
- assert Bs is not None and not is_mx(B, Bs) and weight_block_size(B, Bs) is not None, (
2241
- "a per-tensor scalar As (static activation scale) needs block-scale FP8 weights"
 
 
2242
  )
 
2243
  out = w8a8_block_static_fp8_matmul_grouped(
2244
  A,
2245
  B,
 
24
  from .compat import add_op_namespace_prefix, FP8_DTYPE, NIBBLES_PER_BYTE, compile_time_only_triton_op, compile_time_only_triton_wrap, device_context, get_accelerator_autotuning_configs, sm_count, tl_dtype
25
  from .descriptors import build_grouped_operand_descriptors, rebind_grouped_descriptors, rebind_grouped_mx_descriptors
26
  from .formats import check_activation_format, global_scale_stride, is_per_expert_global, normalize_global_scale, e2m1_as_uint8, expert_weight_shape, is_mx, mx_scale_family, normalize_per_expert_scale, resolve_activation_format, resolve_output_dtype, routed_rows, tokens_per_expert_bucket, ue8m0_as_uint8, validate_dense_operands, weight_block_size, weight_format
27
+ from .quant import fp8_act_quant_block_dynamic, MX_ACT_QUANT, mx_act_quant_grouped, quantize_routed_rows_per_expert, static_expert_act_operands, swizzle_grouped_mx_scales, tensor_wide_act_operands
28
  from .swizzle import swizzled_scale_descriptor
29
  from .mma import block_dynamic_dot, fp8_dot, mx_compute, mx_weight_only_compute, static_dot
30
+ from .scheduling import build_tile_layout, expand_gather_below_parity, expand_regime, build_packed_schedule, load_packed_schedule, prefetch_packed_entry, resolve_grouped_tile, resolve_grouped_tile_packed
31
  from .loading.tiles import (
32
  load_act_block_dynamic,
33
  load_act_mx,
 
41
  weight_tile_ptrs,
42
  )
43
  from .epilogue import acc_init, bias_strides, gemm_epilogue
44
+ from .pruners import PATH_ANCHOR_AXES, fp8_dot_warp_pruner, raw_activation_pointer_pruner, global_scale_warp_spec_pruner, warp_spec_memory_mode_pruner, packed_schedule_scope_pruner, affine_scale_warp_spec_pruner, block_dynamic_grouped_matmul_pruner, block_fits_dim_pruner, block_within_dim_pruner, compose_pruners, descriptor_box_pruner, gate_stacked_tmem_trap_pruner, gated_pointer_weight_warp_spec_pruner, mx_config_pruner, require_moe_dims_aligned, smem_pruner, swizzled_out_bm_pruner, swizzled_scale_config_pruner, swizzled_scales_bm_pruner, warp_spec_compile_guard_pruner
45
 
46
 
47
  @bayesian_autotune(
 
69
  fp8_dot_warp_pruner(),
70
  packed_schedule_scope_pruner(),
71
  block_dynamic_grouped_matmul_pruner(),
72
+ warp_spec_memory_mode_pruner(),
73
  descriptor_box_pruner(),
74
  )
75
  },
 
275
  packed_schedule_scope_pruner(),
276
  block_dynamic_grouped_matmul_pruner(),
277
  descriptor_box_pruner(),
278
+ raw_activation_pointer_pruner(),
279
  )
280
  },
281
  )
 
283
  def w8a8_block_static_fp8_matmul_grouped_kernel(
284
  A, # (num_tokens, K) E4M3 activations (pre-quantized against the static scale by the wrapper)
285
  ADescriptor, # host TMA descriptor over A (rows, K), box (BM, BK); read iff A_MEMORY_MODE != "pointer"
286
+ As, # calibrated (static) activation scale: one value, or one per expert
287
  B, # (num_experts, N, K) FP8 weights; under GATE the (num_experts, 2N, K) gate|up stack
288
  BDescriptor, # host TMA descriptor over B viewed (E, 2N|N, K), box (1, (2|1)*BN, BK); read iff B_MEMORY_MODE != "pointer"
289
  Bs, # (num_experts, N // BLOCK_SIZE_N, K // BLOCK_SIZE_K) weight scales (2N under GATE)
 
304
  stride_b_e,
305
  stride_b_k,
306
  stride_b_n,
307
+ stride_as_e, # 0 = one calibrated scale shared by every expert
308
  stride_bs_e,
309
  stride_bs_k,
310
  stride_bs_n,
 
345
  weight scales apply per-K-tile (plain ``tl.dot`` + software rescale, ``accumulate("static")``),
346
  and the scalar activation scale multiplies the accumulator once after the loop. GATE=False is
347
  the plain grouped GEMM (down projection), bit-identical."""
 
348
  start_pid = tl.program_id(axis=0)
349
  if PACKED_SCHEDULE:
350
  total_m_tiles = load_packed_schedule(Schedule, BLOCK_SIZE_M)
 
393
  gate_s_ptr = Bs + expert_id64 * stride_bs_e + ((2 if GATE else 1) * pid_n) * stride_bs_n
394
  up_s_ptr = gate_s_ptr + stride_bs_n
395
 
396
+ # this tile's expert scale: stride 0 means one calibrated scale for every expert.
397
+ # It feeds the inline quant arm (raw A) and folds back onto the accumulator post-loop.
398
+ a_s_static = tl.load(As + expert_id64.to(tl.int32) * stride_as_e)
399
  acc = acc_init("dot", BLOCK_SIZE_M, (2 if GATE else 1) * BLOCK_SIZE_N, False)
400
  for k in tl.range(0, tl.cdiv(K, BLOCK_SIZE_K), warp_specialize=WARP_SPEC):
401
  a, a_dead = load_act_static(
402
+ a_ptrs, ADescriptor, m_start, k * BLOCK_SIZE_K, row_mask, in_row, a_s_static,
403
  A_MEMORY_MODE, GatherIdx is not None,
404
  )
405
  w, w_s = load_weight_static(
 
471
  warp_spec_compile_guard_pruner(),
472
  descriptor_box_pruner(),
473
  smem_pruner(),
474
+ raw_activation_pointer_pruner(),
475
  )
476
  },
477
  )
 
479
  def w8a8_tensor_dynamic_fp8_matmul_grouped_kernel(
480
  A, # (num_tokens, K) pre-quantized FP8 activations, any row order
481
  ADescriptor, # host TMA descriptor over A (rows, K), box (BM, BK); read iff A_MEMORY_MODE != "pointer"
482
+ As, # per-token activation scales (S,), or a calibrated (static) one: shared, or per expert
483
  B, # (num_experts, N, K) FP8 weights; under GATE the (num_experts, 2N, K) gate|up stack
484
  BDescriptor, # host TMA descriptor over B viewed (E, 2N|N, K), box (1, (2|1)*BN, BK); read iff B_MEMORY_MODE != "pointer"
485
  Bs, # (num_experts, 1, 1) per-tensor weight scales (one scalar covers the gate|up stack)
 
498
  stride_a_m,
499
  stride_a_k,
500
  stride_as_m,
501
+ stride_as_e, # 0 = the scale is per token; 1 = one calibrated scale per expert
502
  stride_b_e,
503
  stride_b_k,
504
  stride_b_n,
 
587
  stride_b_k,
588
  False,
589
  )
590
+ # per token, or the tile's own expert under a calibrated scale (stride_as_m 0); 1.0 on
591
+ # the masked rows so the static arm's in-register divide never sees a zero
592
+ a_s = tl.load(
593
+ As + in_row * stride_as_m + expert_id64.to(tl.int32) * stride_as_e, mask=row_mask, other=1.0
594
+ )
595
  b_s = tl.load(Bs + expert_id64 * stride_bs_e)
596
 
597
  acc = acc_init("dot", BLOCK_SIZE_M, (2 if GATE else 1) * BLOCK_SIZE_N, False)
598
  for k in tl.range(0, tl.cdiv(K, BLOCK_SIZE_K), warp_specialize=WARP_SPEC):
599
+ a, _as = load_act_static(
600
+ a_ptrs, ADescriptor, m_start, k * BLOCK_SIZE_K, row_mask, in_row, a_s[:, None],
601
  A_MEMORY_MODE, GatherIdx is not None,
602
  )
603
  w, _ws = load_weight_plain(
 
676
  smem_pruner(min_sets=2), # two full buffer sets: the grouped pipeline's floor (see smem_pruner)
677
  warp_spec_compile_guard_pruner(),
678
  affine_scale_warp_spec_pruner(),
679
+ warp_spec_memory_mode_pruner(),
680
  )
681
  },
682
  )
 
928
  path_anchor_axes=PATH_ANCHOR_AXES,
929
  prune_configs_by={
930
  "early_config_prune": compose_pruners(
931
+ warp_spec_memory_mode_pruner(weight_only=True),
932
  packed_schedule_scope_pruner(),
933
  # the tail tile slides back in bounds (see the K-loop) — BK only has to fit
934
  block_fits_dim_pruner("K"),
 
1448
 
1449
  A: (S, K) raw bf16/fp16 activations — rows addressed via ``gather_idx``
1450
  B: (num_experts, N, K) FP8 weights; under ``gate`` the (num_experts, 2N, K) gate|up stack
1451
+ As: scalar / (1,) for one calibrated scale, or (num_experts,) for a MoE that
1452
+ calibrates each expert separately
1453
  Bs: (num_experts, N // block_n, K // block_k) per-block weight scales (2N under gate)
1454
  """
1455
  validate_dense_operands(A, B)
 
1476
  )
1477
 
1478
  output_dtype = resolve_output_dtype(output_dtype, A, None)
1479
+ A_q, As, as_stride = static_expert_act_operands(A, As, num_experts)
1480
  bs_u8 = ue8m0_as_uint8(Bs)
 
 
 
1481
  if requant:
1482
  C = A.new_empty(S, N, dtype=FP8_DTYPE)
1483
  # UE8M0 model (ue8m0 weights) -> UE8M0 intermediate scales; the epilogue infers the format
 
1520
  B.stride(0),
1521
  B.stride(2),
1522
  B.stride(1),
1523
+ as_stride,
1524
  bs_u8.stride(0),
1525
  bs_u8.stride(2),
1526
  bs_u8.stride(1),
 
1580
 
1581
  A: (S, K) pre-quantized FP8 activations — rows addressed via ``gather_idx``
1582
  B: (num_experts, N, K) FP8 expert weights; under ``gate`` the (num_experts, 2N, K) stack
1583
+ As: (S,) per-token scales alongside a pre-quantized A, or — on a raw A — the calibrated
1584
+ (static) activation scale: one value, or one per expert
1585
  Bs: (num_experts,) or (num_experts, 1, 1) per-expert weight scales
1586
  expert_start: (num_experts_pow2 + 1,) int32 — cumulative sorted-row starts, S sentinel
1587
  gather_idx: optional (S,) — sorted position -> source row of A; None = A is expert-sorted
 
1605
  # Normalize Bs to (num_experts, 1, 1) — one per-tensor scale (covers the gate|up stack)
1606
  Bs = normalize_per_expert_scale(Bs, num_experts)
1607
 
 
1608
  output_dtype = resolve_output_dtype(output_dtype, A, As)
1609
+ A, As, as_stride_m, as_stride_e = tensor_wide_act_operands(A, As, num_experts)
1610
+ if as_stride_e and expand_regime(S, num_experts):
1611
+ # a per-expert calibrated scale: lay the routed rows out quantized rather than read
1612
+ # them raw through every N-tile — the same copy-amortizes-at-prefill law below
1613
+ A = quantize_routed_rows_per_expert(A, As, gather_idx, expert_start, num_experts)
1614
+ gather_idx = None # spent: the rows are laid out now
1615
+ if A.dtype == FP8_DTYPE:
1616
+ # post-quant: trade the in-kernel gather for one packed-row copy where that wins. Keyed
1617
+ # on A being quantized, which the dynamic quant, a shared calibrated scale and the
1618
+ # per-expert layout above all reach; raw rows (decode) keep the gather.
1619
  A, _, gather_idx = expand_gather_below_parity(A, None, gather_idx, num_experts)
1620
  C = A.new_empty(S, N, dtype=output_dtype)
1621
  num_sms = sm_count(A.device.index)
 
1649
  K,
1650
  A.stride(0),
1651
  A.stride(1),
1652
+ as_stride_m,
1653
+ as_stride_e,
1654
  B.stride(0),
1655
  B.stride(2),
1656
  B.stride(1),
 
2253
  "output_global_scale is the NVFP4 requant second level — it requires quantize_output=True on NVFP4 "
2254
  "(the epilogue would otherwise normalize by it with nothing downstream to compensate)"
2255
  )
2256
+ # a calibrated (static) scale on a raw A (see `tensor_wide_act_operands`): block-scale weights
2257
+ # have a dedicated static kernel, per-tensor ones read the same As on the tensor-wide arm
2258
+ static_act = As is not None and As.ndim <= 1 and As.numel() in (1, B.shape[0]) and A.dtype != FP8_DTYPE
2259
+ if static_act:
2260
+ assert Bs is not None and not is_mx(B, Bs), (
2261
+ "a calibrated (static) activation scale is an FP8 form — MX activations carry a scale "
2262
+ "per group, derived per call"
2263
  )
2264
+ if static_act and weight_block_size(B, Bs) is not None:
2265
  out = w8a8_block_static_fp8_matmul_grouped(
2266
  A,
2267
  B,
build/torch-rocm/loading/tiles.py CHANGED
@@ -268,15 +268,18 @@ def load_act_static(
268
  a_ptrs, a_descriptor, m_off, k_off, value_mask, gather_rows, a_s_static,
269
  A_MEMORY_MODE: tl.constexpr, A_GATHER: tl.constexpr,
270
  ):
271
- """static activation tile. Pre-quantized fp8 A loads as the rows-major MMA lhs; raw bf16/fp16
272
- (inline arm, small M, pointer-only) is quantized against the scalar ``a_s_static``. ``a_s`` = the
273
- values (the static scale is a scalar folded post-loop)."""
274
- if a_ptrs.dtype.element_ty == tl.float8e4nv: # pre-quantized fp8 A (MMA lhs, rows-major)
275
- a = load_grouped_act_tile(
276
- a_ptrs, a_descriptor, m_off, k_off, value_mask, gather_rows, A_MEMORY_MODE, A_GATHER
277
- )
278
- else: # raw bf16/fp16 (inline arm, M<threshold, pointer-only) — quantize vs the static scale
279
- a = (tl.load(a_ptrs).to(tl.float32) / a_s_static).to(tl.float8e4nv)
 
 
 
280
  a_s = a
281
  return a, a_s
282
 
 
268
  a_ptrs, a_descriptor, m_off, k_off, value_mask, gather_rows, a_s_static,
269
  A_MEMORY_MODE: tl.constexpr, A_GATHER: tl.constexpr,
270
  ):
271
+ """static activation tile, through the same tile load in every memory mode: a pre-quantized
272
+ fp8 A is the rows-major MMA lhs as-is, raw bf16/fp16 is quantized in register against
273
+ ``a_s_static``. The raw arm is what a PER-EXPERT scale needs: one row can be routed to
274
+ several experts whose calibrated scales differ, so it has no single pre-quantized form. Its
275
+ own load would have to repeat the tile loader's masking and descriptor arms — a bare
276
+ pointer load reads past the tensor on a grouped tail tile. ``a_s`` = the values (the scale
277
+ is folded post-loop)."""
278
+ a = load_grouped_act_tile(
279
+ a_ptrs, a_descriptor, m_off, k_off, value_mask, gather_rows, A_MEMORY_MODE, A_GATHER
280
+ )
281
+ if a_ptrs.dtype.element_ty != tl.float8e4nv: # raw rows: quantize against the tile's scale
282
+ a = (a.to(tl.float32) / a_s_static).to(tl.float8e4nv)
283
  a_s = a
284
  return a, a_s
285
 
build/torch-rocm/matmul.py CHANGED
@@ -25,7 +25,7 @@ from .compat import add_op_namespace_prefix, FP8_DTYPE, is_sm10x, NIBBLES_PER_BY
25
  from .descriptors import maybe_descriptor, rebind_bd_descriptors, rebind_mx_descriptors, rebind_weight_only_descriptors
26
  from .formats import check_activation_format, global_scale_stride, normalize_global_scale, e2m1_as_uint8, is_mx, mx_scale_family, resolve_activation_format, resolve_output_dtype, ue8m0_as_uint8, validate_dense_2d_operands, weight_format
27
  from .swizzle import swizzle_mx_scales, swizzled_scale_descriptor
28
- from .quant import MX_ACT_QUANT, fp8_act_quant_block_dynamic, fp8_act_quant_tensor_wide, maybe_act_quant, mxfp8_act_quant, nvfp4_act_quant
29
  from .loading.scales import apply_global_scale, mx_2d_scale_ptrs
30
  from .mma import block_dynamic_dot, fp8_dot, mx_compute, mx_weight_only_compute, static_dot
31
  from .loading.tiles import (
@@ -299,7 +299,7 @@ def w8a8_block_dynamic_fp8_matmul_kernel(
299
  @triton.jit
300
  def w8a8_tensor_dynamic_fp8_matmul_kernel(
301
  A, # (M, K) pre-quantized FP8 activations
302
- As, # (M,) per-token activation scales
303
  B, # (N, K) FP8 weights
304
  Bs, # scalar/(1,) per-tensor weight scale
305
  C, # (M, N) output
@@ -1252,6 +1252,7 @@ def w8a8_tensor_dynamic_fp8_matmul(
1252
  A: torch.Tensor,
1253
  B: torch.Tensor,
1254
  Bs: torch.Tensor,
 
1255
  output_dtype: torch.dtype | None = None,
1256
  gate: bool = False,
1257
  act_fn: str = "silu",
@@ -1263,15 +1264,27 @@ def w8a8_tensor_dynamic_fp8_matmul(
1263
  """Tensor-scale FP8 matmul: ``C = A @ B.T``; activations quantized offline per row.
1264
 
1265
  A: (..., K) raw activations, bf16/fp16/fp32 (flattened to (M, K)
1266
- internally) — per-row scales computed via ``fp8_act_quant_tensor_wide(A, K)``.
 
1267
  B: (N, K) FP8 weights — under ``gate`` the ``(2N, K)`` gate|up stack (one per-tensor scale).
1268
  Bs: scalar, (1,), or (1, 1) — single tensor-scale weight scale.
 
 
1269
 
1270
  ``gate`` fuses the gate|up projection into one stacked GEMM + SwiGLU, returning the
1271
  ``[..., N]`` GLU intermediate. Returns a one-element list (mirrors the MX/grouped op).
1272
  """
1273
  validate_dense_2d_operands(A, B)
1274
 
 
 
 
 
 
 
 
 
 
1275
  rows, K = B.shape
1276
  # Under gate|up fusion B is the (2N, K) gate|up stack; N is the per-projection output width.
1277
  N = rows // 2 if gate else rows
@@ -1279,9 +1292,8 @@ def w8a8_tensor_dynamic_fp8_matmul(
1279
 
1280
  assert Bs.numel() == 1, f"Bs must be scalar or (1,), got {tuple(Bs.shape)}"
1281
 
1282
- # Per-row scalar activation scale (one per token).
1283
- qA, As = fp8_act_quant_tensor_wide(A, K)
1284
- As = As.reshape(M)
1285
  Bs = Bs.reshape(1)
1286
 
1287
  C = A.new_empty(A.shape[:-1] + (N,), dtype=output_dtype)
@@ -1307,7 +1319,7 @@ def w8a8_tensor_dynamic_fp8_matmul(
1307
  int(M).bit_length(), # m_bit_length key bucket
1308
  qA.stride(-2),
1309
  qA.stride(-1),
1310
- As.stride(0),
1311
  B.stride(1),
1312
  B.stride(0),
1313
  C.stride(-2),
@@ -1374,15 +1386,8 @@ def mx_dynamic_matmul(
1374
  # direct op calls, prequantized-As callers) and so autograd keeps working: register_autograd
1375
  # hangs on this op and saves INPUTS, so the registered dgrad serves a scaled_mm forward
1376
  # unchanged. Measured 1.4-4.8x over the Triton kernel at every M (B200, 2026-08-27); a None
1377
- # falls through to the Triton launch exactly as before. Skipped under fake/meta propagation
1378
- # (the opaque op's fake impl runs this body with launches disabled — the shape-correct C
1379
- # allocation below is all fake mode needs).
1380
- from . import compat as _compat
1381
-
1382
- if (
1383
- not gate and not quantize_output and bias is None and not simulate_unfused
1384
- and not _compat._SKIP_LAUNCHES_MIRROR
1385
- ):
1386
  _smm = _torch_scaled_mm_2d(
1387
  A, B, As, Bs, activation_format, a_global_scale, b_global_scale,
1388
  resolve_output_dtype(output_dtype, A, None),
@@ -1728,7 +1733,7 @@ def mx_weight_only_matmul_2d(
1728
  return [C]
1729
 
1730
 
1731
- # ── torch scaled_mm fast path (2D dense MX, Blackwell) ────────────────────────
1732
  #
1733
  # cuBLAS's block-scaled GEMM beats our Triton 2D kernel at EVERY measured M on the
1734
  # quantized-activation MX formats (B200, N=18432 K=6144, 200-trial tunes, 2026-08-27):
@@ -1738,6 +1743,14 @@ def mx_weight_only_matmul_2d(
1738
  # quant feeds both, fp32 accumulate). Only the M=1 swap-AB decode dispatch stays ahead
1739
  # (18.4us vs ~33), so the route floor is M >= 2.
1740
  #
 
 
 
 
 
 
 
 
1741
  # Inductor refuses to lower scaled_mm with SWIZZLE_32_4_4 scales ("does not yet support
1742
  # non-trivial swizzles" — repros/scaled_mm_swizzle_compile.py), so the call lives behind an
1743
  # OPAQUE custom op: dynamo captures the op as a leaf and the real scaled_mm runs at runtime,
@@ -1748,7 +1761,7 @@ def mx_weight_only_matmul_2d(
1748
  _SCALED_MM_MIN_M = int(os.environ.get("FINEGRAINED_SCALED_MM_MIN_M", "2"))
1749
 
1750
 
1751
- @torch.library.custom_op(add_op_namespace_prefix("scaled_mm_2d_mx"), mutates_args=())
1752
  def _scaled_mm_2d_op(
1753
  Aq: torch.Tensor,
1754
  As: torch.Tensor,
@@ -1764,6 +1777,13 @@ def _scaled_mm_2d_op(
1764
  cuBLAS fast path too, not just eager ones."""
1765
  F = torch.nn.functional
1766
  ST, SW = F.ScalingType, F.SwizzleType
 
 
 
 
 
 
 
1767
  if fmt == "mxfp8":
1768
  return F.scaled_mm(
1769
  Aq, B.t(),
@@ -1797,6 +1817,38 @@ def _smm_reject(reason):
1797
  return None
1798
 
1799
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1800
  def _torch_scaled_mm_2d(A, B, As, Bs, activation_format, a_global_scale, b_global_scale, out_dtype):
1801
  """The scaled_mm route, or ``None`` to take the Triton ops. Fires only where measured to
1802
  win AND where the semantics are identical: plain ungated GEMM, quantized activations in
@@ -1950,10 +2002,9 @@ def matmul_2d(
1950
  block_n = B.shape[0] // Bs.shape[0] if B.shape[0] % Bs.shape[0] == 0 else block_k
1951
  block_size = [block_n, block_k]
1952
  if block_size is None: # tensor-wide (per-tensor) scale
1953
- assert As is None, "tensor-wide FP8 quantizes A dynamically — no As"
1954
  return _unwrap(
1955
  w8a8_tensor_dynamic_fp8_matmul(
1956
- A, B, Bs, output_dtype, gate, act_fn, swiglu_alpha, swiglu_limit, simulate_unfused, bias=bias)
1957
  )
1958
  # Block-wise FP8: a per-tensor scalar As is the static (calibrated) activation scale; else dynamic.
1959
  if As is not None:
 
25
  from .descriptors import maybe_descriptor, rebind_bd_descriptors, rebind_mx_descriptors, rebind_weight_only_descriptors
26
  from .formats import check_activation_format, global_scale_stride, normalize_global_scale, e2m1_as_uint8, is_mx, mx_scale_family, resolve_activation_format, resolve_output_dtype, ue8m0_as_uint8, validate_dense_2d_operands, weight_format
27
  from .swizzle import swizzle_mx_scales, swizzled_scale_descriptor
28
+ from .quant import MX_ACT_QUANT, fp8_act_quant_block_dynamic, maybe_act_quant, mxfp8_act_quant, nvfp4_act_quant, quantize_rows_static, tensor_wide_act_operands
29
  from .loading.scales import apply_global_scale, mx_2d_scale_ptrs
30
  from .mma import block_dynamic_dot, fp8_dot, mx_compute, mx_weight_only_compute, static_dot
31
  from .loading.tiles import (
 
299
  @triton.jit
300
  def w8a8_tensor_dynamic_fp8_matmul_kernel(
301
  A, # (M, K) pre-quantized FP8 activations
302
+ As, # (M,) per-token activation scales, or one calibrated (static) scale (stride_as_m 0)
303
  B, # (N, K) FP8 weights
304
  Bs, # scalar/(1,) per-tensor weight scale
305
  C, # (M, N) output
 
1252
  A: torch.Tensor,
1253
  B: torch.Tensor,
1254
  Bs: torch.Tensor,
1255
+ As: torch.Tensor | None = None,
1256
  output_dtype: torch.dtype | None = None,
1257
  gate: bool = False,
1258
  act_fn: str = "silu",
 
1264
  """Tensor-scale FP8 matmul: ``C = A @ B.T``; activations quantized offline per row.
1265
 
1266
  A: (..., K) raw activations, bf16/fp16/fp32 (flattened to (M, K)
1267
+ internally) — per-row scales computed via ``fp8_act_quant_tensor_wide(A, K)``, or
1268
+ quantized against ``As`` when a calibrated (static) one is given.
1269
  B: (N, K) FP8 weights — under ``gate`` the ``(2N, K)`` gate|up stack (one per-tensor scale).
1270
  Bs: scalar, (1,), or (1, 1) — single tensor-scale weight scale.
1271
+ As: the calibrated (static) activation scale, one value for the whole matmul; ``None``
1272
+ derives one per row from the data.
1273
 
1274
  ``gate`` fuses the gate|up projection into one stacked GEMM + SwiGLU, returning the
1275
  ``[..., N]`` GLU intermediate. Returns a one-element list (mirrors the MX/grouped op).
1276
  """
1277
  validate_dense_2d_operands(A, B)
1278
 
1279
+ # cuBLAS per-tensor fast path, INSIDE the op for the same reasons as the MX one below:
1280
+ # every caller gets it, and a None falls through to the Triton launch unchanged.
1281
+ if not gate and not simulate_unfused and bias is None:
1282
+ _smm = _torch_scaled_mm_2d_tensor_wise(
1283
+ A, B, As, Bs, None, resolve_output_dtype(output_dtype, A, None)
1284
+ )
1285
+ if _smm is not None:
1286
+ return [_smm]
1287
+
1288
  rows, K = B.shape
1289
  # Under gate|up fusion B is the (2N, K) gate|up stack; N is the per-projection output width.
1290
  N = rows // 2 if gate else rows
 
1292
 
1293
  assert Bs.numel() == 1, f"Bs must be scalar or (1,), got {tuple(Bs.shape)}"
1294
 
1295
+ # one scale per row (derived here), or the calibrated one every row reads (stride 0)
1296
+ qA, As, as_stride_m, _ = tensor_wide_act_operands(A, As)
 
1297
  Bs = Bs.reshape(1)
1298
 
1299
  C = A.new_empty(A.shape[:-1] + (N,), dtype=output_dtype)
 
1319
  int(M).bit_length(), # m_bit_length key bucket
1320
  qA.stride(-2),
1321
  qA.stride(-1),
1322
+ as_stride_m,
1323
  B.stride(1),
1324
  B.stride(0),
1325
  C.stride(-2),
 
1386
  # direct op calls, prequantized-As callers) and so autograd keeps working: register_autograd
1387
  # hangs on this op and saves INPUTS, so the registered dgrad serves a scaled_mm forward
1388
  # unchanged. Measured 1.4-4.8x over the Triton kernel at every M (B200, 2026-08-27); a None
1389
+ # falls through to the Triton launch exactly as before.
1390
+ if not gate and not quantize_output and bias is None and not simulate_unfused:
 
 
 
 
 
 
 
1391
  _smm = _torch_scaled_mm_2d(
1392
  A, B, As, Bs, activation_format, a_global_scale, b_global_scale,
1393
  resolve_output_dtype(output_dtype, A, None),
 
1733
  return [C]
1734
 
1735
 
1736
+ # ── torch scaled_mm fast path (2D dense, Blackwell) ──────────────────────────
1737
  #
1738
  # cuBLAS's block-scaled GEMM beats our Triton 2D kernel at EVERY measured M on the
1739
  # quantized-activation MX formats (B200, N=18432 K=6144, 200-trial tunes, 2026-08-27):
 
1743
  # quant feeds both, fp32 accumulate). Only the M=1 swap-AB decode dispatch stays ahead
1744
  # (18.4us vs ~33), so the route floor is M >= 2.
1745
  #
1746
+ # Per-tensor static FP8 wins at EVERY M, decode included, so that arm has no floor (B200,
1747
+ # N=12288 K=4096, `quantize_rows_static` feeding both arms, cudagraph):
1748
+ # M=1 12.2us vs 14.0 M=8 9.8 vs 13.0 M=64 8.9 vs 12.0 M=512 23.2 vs 34.1
1749
+ # M=8192 320.5us vs 413.0
1750
+ # GEMM alone that is 2807 TFLOP/s against our 1993, and cuBLAS lands within 1.4% of its own
1751
+ # big-square best case (2846) — the 2-CTA + TMA warp-specialized architecture we cannot reach
1752
+ # from Triton, not a tuning gap. Relerr 1.2e-5 against the Triton arm.
1753
+ #
1754
  # Inductor refuses to lower scaled_mm with SWIZZLE_32_4_4 scales ("does not yet support
1755
  # non-trivial swizzles" — repros/scaled_mm_swizzle_compile.py), so the call lives behind an
1756
  # OPAQUE custom op: dynamo captures the op as a leaf and the real scaled_mm runs at runtime,
 
1761
  _SCALED_MM_MIN_M = int(os.environ.get("FINEGRAINED_SCALED_MM_MIN_M", "2"))
1762
 
1763
 
1764
+ @torch.library.custom_op(add_op_namespace_prefix("scaled_mm_2d"), mutates_args=())
1765
  def _scaled_mm_2d_op(
1766
  Aq: torch.Tensor,
1767
  As: torch.Tensor,
 
1777
  cuBLAS fast path too, not just eager ones."""
1778
  F = torch.nn.functional
1779
  ST, SW = F.ScalingType, F.SwizzleType
1780
+ if fmt == "fp8":
1781
+ return F.scaled_mm(
1782
+ Aq, B.t(),
1783
+ As.reshape(()), ST.TensorWise,
1784
+ Bs.reshape(()), ST.TensorWise,
1785
+ output_dtype=torch.bfloat16,
1786
+ )
1787
  if fmt == "mxfp8":
1788
  return F.scaled_mm(
1789
  Aq, B.t(),
 
1817
  return None
1818
 
1819
 
1820
+ def _torch_scaled_mm_2d_tensor_wise(A, B, As, Bs, activation_format, out_dtype):
1821
+ """The scaled_mm route for per-tensor static FP8, or ``None`` to take the Triton ops.
1822
+
1823
+ One fp32 value per operand IS scaled_mm's ``TensorWise``, so this arm needs neither swizzled
1824
+ scales nor a block-aligned K — only a STATIC activation scale, since a dynamic one would cost
1825
+ the amax pass the Triton kernel fuses into its own loop.
1826
+ """
1827
+ if os.environ.get("FINEGRAINED_DISABLE_SCALED_MM"):
1828
+ return _smm_reject("env-disabled")
1829
+ if not is_sm10x():
1830
+ return _smm_reject("not-sm10x")
1831
+ if not (hasattr(torch.nn.functional, "scaled_mm") and hasattr(torch.nn.functional, "ScalingType")):
1832
+ return _smm_reject("no-F.scaled_mm")
1833
+ if activation_format not in (None, "fp8"):
1834
+ return _smm_reject("activation_format")
1835
+ if out_dtype is not torch.bfloat16:
1836
+ return _smm_reject("out-dtype")
1837
+ if As is None or As.numel() != 1:
1838
+ return _smm_reject("not-static-per-tensor")
1839
+ if As.dtype is not torch.float32 or Bs.dtype is not torch.float32:
1840
+ # a UE8M0 container's scalar is an EXPONENT, which TensorWise would read as a multiplier
1841
+ return _smm_reject("non-fp32 scale")
1842
+ K = A.shape[-1]
1843
+ if K % 16 or B.shape[0] % 16:
1844
+ return _smm_reject("shape")
1845
+ M = A.numel() // K
1846
+ As = As.reshape(-1).float()
1847
+ Aq = A.reshape(M, K) if A.dtype == FP8_DTYPE else quantize_rows_static(A.reshape(M, K), As)
1848
+ out = _scaled_mm_2d_op(Aq, As, B, Bs.reshape(-1).float(), None, None, "fp8")
1849
+ return out.reshape(*A.shape[:-1], B.shape[0])
1850
+
1851
+
1852
  def _torch_scaled_mm_2d(A, B, As, Bs, activation_format, a_global_scale, b_global_scale, out_dtype):
1853
  """The scaled_mm route, or ``None`` to take the Triton ops. Fires only where measured to
1854
  win AND where the semantics are identical: plain ungated GEMM, quantized activations in
 
2002
  block_n = B.shape[0] // Bs.shape[0] if B.shape[0] % Bs.shape[0] == 0 else block_k
2003
  block_size = [block_n, block_k]
2004
  if block_size is None: # tensor-wide (per-tensor) scale
 
2005
  return _unwrap(
2006
  w8a8_tensor_dynamic_fp8_matmul(
2007
+ A, B, Bs, As, output_dtype, gate, act_fn, swiglu_alpha, swiglu_limit, simulate_unfused, bias=bias)
2008
  )
2009
  # Block-wise FP8: a per-tensor scalar As is the static (calibrated) activation scale; else dynamic.
2010
  if As is not None:
build/torch-rocm/metadata.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "finegrained-kernels",
3
- "id": "_finegrained_kernels_rocm_1348465",
4
  "version": 0,
5
  "license": "Apache-2.0",
6
  "python-depends": [],
@@ -11,25 +11,25 @@
11
  "algorithm": "sha256",
12
  "files": {
13
  "__init__.py": "43nBPHxybam5ChWZmsz1Ji+pjJ6K6zF5lzwbqPcMwE4=",
14
- "_ops.py": "Ssu+B/vlZZOXRh+hKJNWiv3V59Z7J6Wf0w/+9c7uT+M=",
15
  "backward.py": "ri9lKPhb9SmkYBvUS8KvuORiW5PcKViNamFNVWxb69g=",
16
- "batched.py": "cSkV7Mnv6kK+k5MANm8f6ZFu3Jdbh0tAz9u9t4e+R3c=",
17
- "bayesian_autotuner.py": "DPpIATTls8aY7n12oKggpuAOvC3RlugYl1Wid1t5J+U=",
18
- "compat.py": "yj7b7kCAOrq3j1Y3WxSb/YKoVcPRFfkeXmMj7dkqQoc=",
19
  "descriptors.py": "4acRMazGPFlbydNt1Hbo8RpAqb3fOtnOQBv5yb4onqw=",
20
  "epilogue.py": "tGALTI/VMgO3Mz0sQ6H6wuRZpO4UKqM1/4h7d81SnVM=",
21
  "finegrained_kernels/__init__.py": "DFYPlrhXwYjEqCl/8n0SmWGZV8NFml5DPhMjKfv98GY=",
22
  "formats.py": "V6F0PXdNhXQn/UBz0/frbvcp5Lgterh053XDrvQ/Cpw=",
23
- "grouped.py": "PXhZU9Rn6UIIiQtR6IUmI060Mank5GMuooTWapRvya8=",
24
  "loading/__init__.py": "so7kYmlVoiASfoOzemdU5fBTiP8yA85gVPNQExIq/KU=",
25
  "loading/scales.py": "8HJpflmImEXYJquA6hVLStUUEecvExAxBEYVb3qPPWE=",
26
- "loading/tiles.py": "hvFIDikI1k0+vV1QOjqdl1bovI/JEElfNwLWg7VC3ZM=",
27
- "matmul.py": "qLV1KsPyCeAq3O0swnk38GUOT+lS12WxUM8bQZ9DF9A=",
28
  "mma.py": "1fg2an7Dc8Vd6qpSmq+Q0781A6sNilOJHU0PqEbouT8=",
29
- "moe.py": "rpiYMM+kzTtNuxXFUrQARLO9kLHcEgKR8rs/dZ0wn/g=",
30
  "norm.py": "VvE6GUE0Oztf7PDcviLCcB4fhnXt7/zgdE+aApv6DBM=",
31
- "pruners.py": "v03bXbvsZp1znVIj/eiR7LAWctU66DJlJr51wn8Pheg=",
32
- "quant.py": "PZhzgOHSDH8lLV6zvrv040B4g9SgFE6Eq6bc8XxfPY8=",
33
  "scheduling.py": "UGgQ5uVdm4cWPnH+4w6Ffou/PSX3iUarMeRo2E04BhY=",
34
  "swizzle.py": "ueIJ9FsO8QZhVt3nMFlA+8b+U9/ZHabHFlhufsONV04="
35
  }
@@ -41,7 +41,7 @@
41
  "dirty": false
42
  },
43
  "kernel": {
44
- "sha": "1348465af00f2d4c596d37bdb262ae0875dc1d77",
45
  "dirty": false
46
  }
47
  }
 
1
  {
2
  "name": "finegrained-kernels",
3
+ "id": "_finegrained_kernels_rocm_b0c6080",
4
  "version": 0,
5
  "license": "Apache-2.0",
6
  "python-depends": [],
 
11
  "algorithm": "sha256",
12
  "files": {
13
  "__init__.py": "43nBPHxybam5ChWZmsz1Ji+pjJ6K6zF5lzwbqPcMwE4=",
14
+ "_ops.py": "EL67XyYKmJCWdvLz7dUF+v4b0yedbiW265F3nOEoZXw=",
15
  "backward.py": "ri9lKPhb9SmkYBvUS8KvuORiW5PcKViNamFNVWxb69g=",
16
+ "batched.py": "683kgUc1fTmvPZPwNha4+VRQMEdZWAULFtDH5/rs8pg=",
17
+ "bayesian_autotuner.py": "GoXX4MLMduf7L55t66joOgfMjRIt4NCQt3aGdon9qJU=",
18
+ "compat.py": "QUGxtD0TogkxhHX8nNNFtwOyjd8ECN72CUGefiwkL/M=",
19
  "descriptors.py": "4acRMazGPFlbydNt1Hbo8RpAqb3fOtnOQBv5yb4onqw=",
20
  "epilogue.py": "tGALTI/VMgO3Mz0sQ6H6wuRZpO4UKqM1/4h7d81SnVM=",
21
  "finegrained_kernels/__init__.py": "DFYPlrhXwYjEqCl/8n0SmWGZV8NFml5DPhMjKfv98GY=",
22
  "formats.py": "V6F0PXdNhXQn/UBz0/frbvcp5Lgterh053XDrvQ/Cpw=",
23
+ "grouped.py": "14Zfza8MI033MXY2JFXVDfaNVbKHPqiYletmBazGwMo=",
24
  "loading/__init__.py": "so7kYmlVoiASfoOzemdU5fBTiP8yA85gVPNQExIq/KU=",
25
  "loading/scales.py": "8HJpflmImEXYJquA6hVLStUUEecvExAxBEYVb3qPPWE=",
26
+ "loading/tiles.py": "tcStORrCdZbFeoG4AXIE3YznJXGt8rIIBHnyNumZMCE=",
27
+ "matmul.py": "Wfmnyoez0xwplElVEd4zpzc0KQ75rr3m09eww6aOFBE=",
28
  "mma.py": "1fg2an7Dc8Vd6qpSmq+Q0781A6sNilOJHU0PqEbouT8=",
29
+ "moe.py": "Mj+n8CNqL0mOQ5xmqpopisUWGwKoayanfaAymXiGhRk=",
30
  "norm.py": "VvE6GUE0Oztf7PDcviLCcB4fhnXt7/zgdE+aApv6DBM=",
31
+ "pruners.py": "/uWNcx3KqQwJGG1FAeJOSY7Jha9PwBsczVfKYQb5GEI=",
32
+ "quant.py": "Wyw3yCpgybFHswP5UEsl13BQ1Hqgm61Z4/XL2TGhvfk=",
33
  "scheduling.py": "UGgQ5uVdm4cWPnH+4w6Ffou/PSX3iUarMeRo2E04BhY=",
34
  "swizzle.py": "ueIJ9FsO8QZhVt3nMFlA+8b+U9/ZHabHFlhufsONV04="
35
  }
 
41
  "dirty": false
42
  },
43
  "kernel": {
44
+ "sha": "b0c60801e54949c92345c0a86a32f36a1fd2885f",
45
  "dirty": false
46
  }
47
  }
build/torch-rocm/metadata.json.sigstore CHANGED
@@ -1 +1 @@
1
- {"mediaType":"application/vnd.dev.sigstore.bundle.v0.3+json", "verificationMaterial":{"certificate":{"rawBytes":"MIIHTDCCBtKgAwIBAgIUKwy7DRLjNl12oimua1LRYxRBXBQwCgYIKoZIzj0EAwMwNzEVMBMGA1UEChMMc2lnc3RvcmUuZGV2MR4wHAYDVQQDExVzaWdzdG9yZS1pbnRlcm1lZGlhdGUwHhcNMjYwOTE2MTU0MjA3WhcNMjYwOTE2MTU1MjA3WjAAMFkwEwYHKoZIzj0CAQYIKoZIzj0DAQcDQgAEkX51E0716VRnV+N7XqZPQSuT+8+WFqVZ8jyLnA6PxNlNCIgr4O1c3VI4np8PMZ+9k0EozfVf3XbIiPsqY9vU/aOCBfEwggXtMA4GA1UdDwEB/wQEAwIHgDATBgNVHSUEDDAKBggrBgEFBQcDAzAdBgNVHQ4EFgQUo1O1wLYeWm7XlvpitaZXgdwuluUwHwYDVR0jBBgwFoAU39Ppz1YkEZb5qNjpKFWixi4YZD8wawYDVR0RAQH/BGEwX4ZdaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5Ly5naXRodWIvd29ya2Zsb3dzL2J1aWxkLnlhbWxAcmVmcy9oZWFkcy9tYWluMDkGCisGAQQBg78wAQEEK2h0dHBzOi8vdG9rZW4uYWN0aW9ucy5naXRodWJ1c2VyY29udGVudC5jb20wHwYKKwYBBAGDvzABAgQRd29ya2Zsb3dfZGlzcGF0Y2gwNgYKKwYBBAGDvzABAwQoMTM0ODQ2NWFmMDBmMmQ0YzU5NmQzN2JkYjI2MmFlMDg3NWRjMWQ3NzATBgorBgEEAYO/MAEEBAVCdWlsZDArBgorBgEEAYO/MAEFBB1odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eTAdBgorBgEEAYO/MAEGBA9yZWZzL2hlYWRzL21haW4wOwYKKwYBBAGDvzABCAQtDCtodHRwczovL3Rva2VuLmFjdGlvbnMuZ2l0aHVidXNlcmNvbnRlbnQuY29tMG0GCisGAQQBg78wAQkEXwxdaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5Ly5naXRodWIvd29ya2Zsb3dzL2J1aWxkLnlhbWxAcmVmcy9oZWFkcy9tYWluMDgGCisGAQQBg78wAQoEKgwoMTM0ODQ2NWFmMDBmMmQ0YzU5NmQzN2JkYjI2MmFlMDg3NWRjMWQ3NzAbBgorBgEEAYO/MAELBA0MC3NlbGYtaG9zdGVkMEAGCisGAQQBg78wAQwEMgwwaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5MDgGCisGAQQBg78wAQ0EKgwoMTM0ODQ2NWFmMDBmMmQ0YzU5NmQzN2JkYjI2MmFlMDg3NWRjMWQ3NzAfBgorBgEEAYO/MAEOBBEMD3JlZnMvaGVhZHMvbWFpbjAaBgorBgEEAYO/MAEPBAwMCjEwNzE0NzU1MjkwLgYKKwYBBAGDvzABEAQgDB5odHRwczovL2dpdGh1Yi5jb20vaHVnZ2luZ2ZhY2UwGAYKKwYBBAGDvzABEQQKDAgyNTcyMDc0MzBtBgorBgEEAYO/MAESBF8MXWh0dHBzOi8vZ2l0aHViLmNvbS9odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eS8uZ2l0aHViL3dvcmtmbG93cy9idWlsZC55YW1sQHJlZnMvaGVhZHMvbWFpbjA4BgorBgEEAYO/MAETBCoMKDEzNDg0NjVhZjAwZjJkNGM1OTZkMzdiZGIyNjJhZTA4NzVkYzFkNzcwIQYKKwYBBAGDvzABFAQTDBF3b3JrZmxvd19kaXNwYXRjaDBkBgorBgEEAYO/MAEVBFYMVGh0dHBzOi8vZ2l0aHViLmNvbS9odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eS9hY3Rpb25zL3J1bnMvMzUxMTY2NDkyOTcvYXR0ZW1wdHMvMTAWBgorBgEEAYO/MAEWBAgMBnB1YmxpYzBGBgorBgEEAYO/MAEYBDgMNnJlcG86aHVnZ2luZ2ZhY2Uva2VybmVscy1jb21tdW5pdHk6cmVmOnJlZnMvaGVhZHMvbWFpbjCBiwYKKwYBBAHWeQIEAgR9BHsAeQB3AN09MGrGxxEyYxkeHJlnNwKiSl643jyt/4eKcoAvKe6OAAABoKrh6swAAAQDAEgwRgIhANX/7k2eKwwKjy8ZxKOpmVkg1xAd5OZZtFP+JoYLm5xdAiEA541PEWCxlJ8vBCctUblYJ8CwahUQrQHAB1az4ZQBRvAwCgYIKoZIzj0EAwMDaAAwZQIxAOrFZmrON23FouSKXOvH3VhA3MDrJ0zVLduijdnHKbEB6XAwuqVcNzLd8e+SystARQIwOdpWPy0GsRlCrzjovhZ/iOd3XOZtYhCHbwuhADV5U8LyW0FHy/Jl3g8mCcNkM+hY"}, "tlogEntries":[{"logIndex":"2863678766", "logId":{"keyId":"wNI9atQGlz+VWfO6LRygH4QUfY/8W4RFwiT5i5WRgB0="}, "kindVersion":{"kind":"hashedrekord", "version":"0.0.1"}, "integratedTime":"1789573327", "inclusionPromise":{"signedEntryTimestamp":"MEUCIQDOUMA7NduNwOul83R0X3Ocp2ybE707LUZ6Z7MxIDEDHwIgY3TKJTPvt660IBYW6CxkGv0W7iyJWsMQ6DE9TAVfu2A="}, "inclusionProof":{"logIndex":"2741774504", "rootHash":"h7JnsSI02oktCuUf2hPr37WQeR1av+o8FGUwr34KmHo=", "treeSize":"2741774526", "hashes":["1O4OieGYomUIdapwcW0fU0+2+02pFm1aEo0ti4sg08s=", "mcj/XMUVpOVz+S5LUKHconnoT5H4349XaBWB7RFeOkI=", "Wdf+QLOD4triMw+JPnM0auon4nMf4ww9lucCEay94fU=", "6OX+tq7oFddP7/kDlt87XDsQ0WlCJnt1SdJUYyCdR+8=", "B5h5p0vseS2PGxobaWllh1w2JdneoJtoZSCxox8DQMc=", "66RI6MJezQM4018f5MMtQln23ah4XGjWHARyTP4OGH4=", "Ojxn+u3l9okV6fn1qhDIlj0xB1NlpdMwwfQZrYDTTDA=", "dNQ0Q6e05QOPvF+9aMH7hKeNMC+sdDn/+tAhAdOJnpQ=", "cuse2fF4Cl9QdYLPPxCWd9gSEI01WyuvnM3KB1Ah78Y=", "g+kZXC8JlqtkZn+3BtMPZOIjPv/1QgiVvlH7tpuF5rU=", "z7hfW7IyWq58tcWTZuli3Brmb7xA7392X7mpl47fnCc=", "/DBUF30Sdk4LDt4XUgHCGSr10fT7nDkSitYbygYwfNg=", "vC+yZlSpf9tVD+3MdNtHCe71eNNHHi6gplgltL/tmDc=", "gGfwuhE97xnhmxu6exffPtS05ABS4/ldyg66sFcxSRE=", "jhevWFOJ2cWU6SvzMnAe08zHGcVquRlP7gAed39+wI4=", "qxzHanAzz57SDdmJe0B7bJK72NTIBbwEMGBBKvDOROw=", "xH/DCseLHr9eKoYT8qsORZK7zVdEGYWHuVtsVrD95wY="], "checkpoint":{"envelope":"rekor.sigstore.dev - 1193050959916656506\n2741774526\nh7JnsSI02oktCuUf2hPr37WQeR1av+o8FGUwr34KmHo=\n\n— rekor.sigstore.dev wNI9ajBGAiEA1L58wqYgcdSBeA6sGCi5zk9nSn7U/bA/v0obbkuIovICIQDm2hhhrVGvjLnHrfJ4aGYqS45uNAgpXJMtOVd90NuZzQ==\n"}}, "canonicalizedBody":"eyJhcGlWZXJzaW9uIjoiMC4wLjEiLCJraW5kIjoiaGFzaGVkcmVrb3JkIiwic3BlYyI6eyJkYXRhIjp7Imhhc2giOnsiYWxnb3JpdGhtIjoic2hhMjU2IiwidmFsdWUiOiJkZmI0MDlkYzcyN2RjMWM1M2I1MGJjZGZjMzU0NmE4NmZjNWNjOWVlY2UzOTVhZmJkZmEzZWJhOWIwNmRjZDY5In19LCJzaWduYXR1cmUiOnsiY29udGVudCI6Ik1FVUNJUUNjRFZZNFd2Ylhob3pHcDFFSTZSdWxNUTJOZ295M2w1OWpxUWxtcVpXem93SWdKaUYxQ0RhWUl0Tm5icC84WDdMd05kQW03bVVycW93ZWJaL2RMall1dzRNPSIsInB1YmxpY0tleSI6eyJjb250ZW50IjoiTFMwdExTMUNSVWRKVGlCRFJWSlVTVVpKUTBGVVJTMHRMUzB0Q2sxSlNVaFVSRU5EUW5STFowRjNTVUpCWjBsVlMzZDVOMFJTVEdwT2JERXliMmx0ZFdFeFRGSlplRkpDV0VKUmQwTm5XVWxMYjFwSmVtb3dSVUYzVFhjS1RucEZWazFDVFVkQk1WVkZRMmhOVFdNeWJHNWpNMUoyWTIxVmRWcEhWakpOVWpSM1NFRlpSRlpSVVVSRmVGWjZZVmRrZW1SSE9YbGFVekZ3WW01U2JBcGpiVEZzV2tkc2FHUkhWWGRJYUdOT1RXcFpkMDlVUlRKTlZGVXdUV3BCTTFkb1kwNU5hbGwzVDFSRk1rMVVWVEZOYWtFelYycEJRVTFHYTNkRmQxbElDa3R2V2tsNmFqQkRRVkZaU1V0dldrbDZhakJFUVZGalJGRm5RVVZyV0RVeFJUQTNNVFpXVW01V0swNDNXSEZhVUZGVGRWUXJPQ3RYUm5GV1dqaHFlVXdLYmtFMlVIaE9iRTVEU1dkeU5FOHhZek5XU1RSdWNEaFFUVm9yT1dzd1JXOTZabFptTTFoaVNXbFFjM0ZaT1haVkwyRlBRMEptUlhkbloxaDBUVUUwUndwQk1WVmtSSGRGUWk5M1VVVkJkMGxJWjBSQlZFSm5UbFpJVTFWRlJFUkJTMEpuWjNKQ1owVkdRbEZqUkVGNlFXUkNaMDVXU0ZFMFJVWm5VVlZ2TVU4eENuZE1XV1ZYYlRkWWJIWndhWFJoV2xoblpIZDFiSFZWZDBoM1dVUldVakJxUWtKbmQwWnZRVlV6T1ZCd2VqRlphMFZhWWpWeFRtcHdTMFpYYVhocE5Ga0tXa1E0ZDJGM1dVUldVakJTUVZGSUwwSkhSWGRZTkZwa1lVaFNNR05JVFRaTWVUbHVZVmhTYjJSWFNYVlpNamwwVERKb01Wb3laSEJpYldSdFdWZE9iQXBNTW5Sc1kyMDFiR0pJVFhSWk1qbDBZbGhXZFdGWVVqVk1lVFZ1WVZoU2IyUlhTWFprTWpsNVlUSmFjMkl6WkhwTU1rb3hZVmQ0YTB4dWJHaGlWM2hCQ21OdFZtMWplVGx2V2xkR2EyTjVPWFJaVjJ4MVRVUnJSME5wYzBkQlVWRkNaemM0ZDBGUlJVVkxNbWd3WkVoQ2VrOXBPSFprUnpseVdsYzBkVmxYVGpBS1lWYzVkV041Tlc1aFdGSnZaRmRLTVdNeVZubFpNamwxWkVkV2RXUkROV3BpTWpCM1NIZFpTMHQzV1VKQ1FVZEVkbnBCUWtGblVWSmtNamw1WVRKYWN3cGlNMlJtV2tkc2VtTkhSakJaTW1kM1RtZFpTMHQzV1VKQ1FVZEVkbnBCUWtGM1VXOU5WRTB3VDBSUk1rNVhSbTFOUkVKdFRXMVJNRmw2VlRWT2JWRjZDazR5U210WmFra3lUVzFHYkUxRVp6Tk9WMUpxVFZkUk0wNTZRVlJDWjI5eVFtZEZSVUZaVHk5TlFVVkZRa0ZXUTJSWGJITmFSRUZ5UW1kdmNrSm5SVVVLUVZsUEwwMUJSVVpDUWpGdlpGZGtibUZYTlc1YWJVWnFXbE01Y2xwWVNuVmFWM2g2VEZkT2RtSlhNVEZpYld3d1pWUkJaRUpuYjNKQ1owVkZRVmxQTHdwTlFVVkhRa0U1ZVZwWFducE1NbWhzV1ZkU2Vrd3lNV2hoVnpSM1QzZFpTMHQzV1VKQ1FVZEVkbnBCUWtOQlVYUkVRM1J2WkVoU2QyTjZiM1pNTTFKMkNtRXlWblZNYlVacVpFZHNkbUp1VFhWYU1td3dZVWhXYVdSWVRteGpiVTUyWW01U2JHSnVVWFZaTWpsMFRVY3dSME5wYzBkQlVWRkNaemM0ZDBGUmEwVUtXSGQ0WkdGSVVqQmpTRTAyVEhrNWJtRllVbTlrVjBsMVdUSTVkRXd5YURGYU1tUndZbTFrYlZsWFRteE1NblJzWTIwMWJHSklUWFJaTWpsMFlsaFdkUXBoV0ZJMVRIazFibUZZVW05a1YwbDJaREk1ZVdFeVduTmlNMlI2VERKS01XRlhlR3RNYm14b1lsZDRRV050Vm0xamVUbHZXbGRHYTJONU9YUlpWMngxQ2sxRVowZERhWE5IUVZGUlFtYzNPSGRCVVc5RlMyZDNiMDFVVFRCUFJGRXlUbGRHYlUxRVFtMU5iVkV3V1hwVk5VNXRVWHBPTWtwcldXcEpNazF0Um13S1RVUm5NMDVYVW1wTlYxRXpUbnBCWWtKbmIzSkNaMFZGUVZsUEwwMUJSVXhDUVRCTlF6Tk9iR0pIV1hSaFJ6bDZaRWRXYTAxRlFVZERhWE5IUVZGUlFncG5OemgzUVZGM1JVMW5kM2RoU0ZJd1kwaE5Oa3g1T1c1aFdGSnZaRmRKZFZreU9YUk1NbWd4V2pKa2NHSnRaRzFaVjA1c1RESjBiR050Tld4aVNFMTBDbGt5T1hSaVdGWjFZVmhTTlUxRVowZERhWE5IUVZGUlFtYzNPSGRCVVRCRlMyZDNiMDFVVFRCUFJGRXlUbGRHYlUxRVFtMU5iVkV3V1hwVk5VNXRVWG9LVGpKS2ExbHFTVEpOYlVac1RVUm5NMDVYVW1wTlYxRXpUbnBCWmtKbmIzSkNaMFZGUVZsUEwwMUJSVTlDUWtWTlJETktiRnB1VFhaaFIxWm9Xa2hOZGdwaVYwWndZbXBCWVVKbmIzSkNaMFZGUVZsUEwwMUJSVkJDUVhkTlEycEZkMDU2UlRCT2VsVXhUV3ByZDB4bldVdExkMWxDUWtGSFJIWjZRVUpGUVZGbkNrUkNOVzlrU0ZKM1kzcHZka3d5WkhCa1IyZ3hXV2sxYW1JeU1IWmhTRlp1V2pKc2RWb3lXbWhaTWxWM1IwRlpTMHQzV1VKQ1FVZEVkbnBCUWtWUlVVc0tSRUZuZVU1VVkzbE5SR013VFhwQ2RFSm5iM0pDWjBWRlFWbFBMMDFCUlZOQ1JqaE5XRmRvTUdSSVFucFBhVGgyV2pKc01HRklWbWxNYlU1MllsTTVid3BrVjJSdVlWYzFibHB0Um1wYVV6bHlXbGhLZFZwWGVIcE1WMDUyWWxjeE1XSnRiREJsVXpoMVdqSnNNR0ZJVm1sTU0yUjJZMjEwYldKSE9UTmplVGxwQ21SWGJITmFRelUxV1ZjeGMxRklTbXhhYmsxMllVZFdhRnBJVFhaaVYwWndZbXBCTkVKbmIzSkNaMFZGUVZsUEwwMUJSVlJDUTI5TlMwUkZlazVFWnpBS1RtcFdhRnBxUVhkYWFrcHJUa2ROTVU5VVdtdE5lbVJwV2tkSmVVNXFTbWhhVkVFMFRucFdhMWw2Um10T2VtTjNTVkZaUzB0M1dVSkNRVWRFZG5wQlFncEdRVkZVUkVKR00ySXpTbkphYlhoMlpERTVhMkZZVG5kWldGSnFZVVJDYTBKbmIzSkNaMFZGUVZsUEwwMUJSVlpDUmxsTlZrZG9NR1JJUW5wUGFUaDJDbG95YkRCaFNGWnBURzFPZG1KVE9XOWtWMlJ1WVZjMWJscHRSbXBhVXpseVdsaEtkVnBYZUhwTVYwNTJZbGN4TVdKdGJEQmxVemxvV1ROU2NHSXlOWG9LVEROS01XSnVUWFpOZWxWNFRWUlpNazVFYTNsUFZHTjJXVmhTTUZwWE1YZGtTRTEyVFZSQlYwSm5iM0pDWjBWRlFWbFBMMDFCUlZkQ1FXZE5RbTVDTVFwWmJYaHdXWHBDUjBKbmIzSkNaMFZGUVZsUEwwMUJSVmxDUkdkTlRtNUtiR05IT0RaaFNGWnVXakpzZFZveVdtaFpNbFYyWVRKV2VXSnRWbk5qZVRGcUNtSXlNWFJrVnpWd1pFaHJObU50Vm0xUGJrcHNXbTVOZG1GSFZtaGFTRTEyWWxkR2NHSnFRMEpwZDFsTFMzZFpRa0pCU0ZkbFVVbEZRV2RTT1VKSWMwRUtaVkZDTTBGT01EbE5SM0pIZUhoRmVWbDRhMlZJU214dVRuZExhVk5zTmpRemFubDBMelJsUzJOdlFYWkxaVFpQUVVGQlFtOUxjbWcyYzNkQlFVRlJSQXBCUldkM1VtZEphRUZPV0M4M2F6SmxTM2QzUzJwNU9GcDRTMDl3YlZaclp6RjRRV1ExVDFwYWRFWlFLMHB2V1V4dE5YaGtRV2xGUVRVME1WQkZWME40Q214S09IWkNRMk4wVldKc1dVbzRRM2RoYUZWUmNsRklRVUl4WVhvMFdsRkNVblpCZDBObldVbExiMXBKZW1vd1JVRjNUVVJoUVVGM1dsRkplRUZQY2tZS1dtMXlUMDR5TTBadmRWTkxXRTkyU0ROV2FFRXpUVVJ5U2pCNlZreGtkV2xxWkc1SVMySkZRalpZUVhkMWNWWmpUbnBNWkRobEsxTjVjM1JCVWxGSmR3cFBaSEJYVUhrd1IzTlNiRU55ZW1wdmRtaGFMMmxQWkROWVQxcDBXV2hEU0dKM2RXaEJSRlkxVlRoTWVWY3dSa2g1TDBwc00yYzRiVU5qVG10TksyaFpDaTB0TFMwdFJVNUVJRU5GVWxSSlJrbERRVlJGTFMwdExTMEsifX19fQ=="}], "timestampVerificationData":{"rfc3161Timestamps":[{"signedTimestamp":"MIICyDADAgEAMIICvwYJKoZIhvcNAQcCoIICsDCCAqwCAQMxDTALBglghkgBZQMEAgEwgbcGCyqGSIb3DQEJEAEEoIGnBIGkMIGhAgEBBgkrBgEEAYO/MAIwMTANBglghkgBZQMEAgEFAAQgEEWXb1aAU9uWYKBG9k5yRgiCru6MaWJqEJ7eotv6QLsCFAVQ7WjJb8AfpOlPrlg4vundjMeXGA8yMDI2MDkxNjE1NDIwN1owAwIBAaAypDAwLjEVMBMGA1UEChMMc2lnc3RvcmUuZGV2MRUwEwYDVQQDEwxzaWdzdG9yZS10c2GgADGCAdowggHWAgEBMFEwOTEVMBMGA1UEChMMc2lnc3RvcmUuZGV2MSAwHgYDVQQDExdzaWdzdG9yZS10c2Etc2VsZnNpZ25lZAIUOhNULwyQYe68wUMvy4qOiyojiwwwCwYJYIZIAWUDBAIBoIH8MBoGCSqGSIb3DQEJAzENBgsqhkiG9w0BCRABBDAcBgkqhkiG9w0BCQUxDxcNMjYwOTE2MTU0MjA3WjAvBgkqhkiG9w0BCQQxIgQgnhWF5DW3tegD61eDeAX0KIpasZoCwslhIsaRZtJyxHwwgY4GCyqGSIb3DQEJEAIvMX8wfTB7MHkEIIX5J7wHq2LKw7RDVsEO/IGyxog/2nq55thw2dE6zQW3MFUwPaQ7MDkxFTATBgNVBAoTDHNpZ3N0b3JlLmRldjEgMB4GA1UEAxMXc2lnc3RvcmUtdHNhLXNlbGZzaWduZWQCFDoTVC8MkGHuvMFDL8uKjosqI4sMMAoGCCqGSM49BAMCBGYwZAIwQ8JSNadEl6jXNYx0SryoCfLxVtpQES91e45sgot7605M9yXk+nFgOfZ5qRSu1fJdAjB3nmagQ7GQmID30bAxY5/nRpO4VFWhFXqvRv5R0UcT17jG0LE4x4TOt9vkd1x9JHE="}]}}, "messageSignature":{"messageDigest":{"algorithm":"SHA2_256", "digest":"37QJ3HJ9wcU7ULzfw1Rqhvxcye7OOVr736PrqbBtzWk="}, "signature":"MEUCIQCcDVY4WvbXhozGp1EI6RulMQ2Ngoy3l59jqQlmqZWzowIgJiF1CDaYItNnbp/8X7LwNdAm7mUrqowebZ/dLjYuw4M="}}
 
1
+ {"mediaType":"application/vnd.dev.sigstore.bundle.v0.3+json", "verificationMaterial":{"certificate":{"rawBytes":"MIIHSzCCBtGgAwIBAgIUZ81snzjoHnJyxgJ4p7A9yWtxEb4wCgYIKoZIzj0EAwMwNzEVMBMGA1UEChMMc2lnc3RvcmUuZGV2MR4wHAYDVQQDExVzaWdzdG9yZS1pbnRlcm1lZGlhdGUwHhcNMjYwOTIxMTYyOTUxWhcNMjYwOTIxMTYzOTUxWjAAMFkwEwYHKoZIzj0CAQYIKoZIzj0DAQcDQgAE7W5TLpJgVdBM2F/NFhNhCgI3KCd12XeXtQ2rv9vspmSCzZ4GQL4edPNf8mvnZpN2AZvlzrluj3wmXKXB/bpyLKOCBfAwggXsMA4GA1UdDwEB/wQEAwIHgDATBgNVHSUEDDAKBggrBgEFBQcDAzAdBgNVHQ4EFgQU+ywNib8rPeO1AnFU7Bs2vkwfuoAwHwYDVR0jBBgwFoAU39Ppz1YkEZb5qNjpKFWixi4YZD8wawYDVR0RAQH/BGEwX4ZdaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5Ly5naXRodWIvd29ya2Zsb3dzL2J1aWxkLnlhbWxAcmVmcy9oZWFkcy9tYWluMDkGCisGAQQBg78wAQEEK2h0dHBzOi8vdG9rZW4uYWN0aW9ucy5naXRodWJ1c2VyY29udGVudC5jb20wHwYKKwYBBAGDvzABAgQRd29ya2Zsb3dfZGlzcGF0Y2gwNgYKKwYBBAGDvzABAwQoYjBjNjA4MDFlNTQ5NDljOTIzNDVjMGE4NmEzMmYzNmExZmQyODg1ZjATBgorBgEEAYO/MAEEBAVCdWlsZDArBgorBgEEAYO/MAEFBB1odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eTAdBgorBgEEAYO/MAEGBA9yZWZzL2hlYWRzL21haW4wOwYKKwYBBAGDvzABCAQtDCtodHRwczovL3Rva2VuLmFjdGlvbnMuZ2l0aHVidXNlcmNvbnRlbnQuY29tMG0GCisGAQQBg78wAQkEXwxdaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5Ly5naXRodWIvd29ya2Zsb3dzL2J1aWxkLnlhbWxAcmVmcy9oZWFkcy9tYWluMDgGCisGAQQBg78wAQoEKgwoYjBjNjA4MDFlNTQ5NDljOTIzNDVjMGE4NmEzMmYzNmExZmQyODg1ZjAbBgorBgEEAYO/MAELBA0MC3NlbGYtaG9zdGVkMEAGCisGAQQBg78wAQwEMgwwaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5MDgGCisGAQQBg78wAQ0EKgwoYjBjNjA4MDFlNTQ5NDljOTIzNDVjMGE4NmEzMmYzNmExZmQyODg1ZjAfBgorBgEEAYO/MAEOBBEMD3JlZnMvaGVhZHMvbWFpbjAaBgorBgEEAYO/MAEPBAwMCjEwNzE0NzU1MjkwLgYKKwYBBAGDvzABEAQgDB5odHRwczovL2dpdGh1Yi5jb20vaHVnZ2luZ2ZhY2UwGAYKKwYBBAGDvzABEQQKDAgyNTcyMDc0MzBtBgorBgEEAYO/MAESBF8MXWh0dHBzOi8vZ2l0aHViLmNvbS9odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eS8uZ2l0aHViL3dvcmtmbG93cy9idWlsZC55YW1sQHJlZnMvaGVhZHMvbWFpbjA4BgorBgEEAYO/MAETBCoMKGIwYzYwODAxZTU0OTQ5YzkyMzQ1YzBhODZhMzJmMzZhMWZkMjg4NWYwIQYKKwYBBAGDvzABFAQTDBF3b3JrZmxvd19kaXNwYXRjaDBkBgorBgEEAYO/MAEVBFYMVGh0dHBzOi8vZ2l0aHViLmNvbS9odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eS9hY3Rpb25zL3J1bnMvMzU2MjUzMjM5MTcvYXR0ZW1wdHMvMTAWBgorBgEEAYO/MAEWBAgMBnB1YmxpYzBGBgorBgEEAYO/MAEYBDgMNnJlcG86aHVnZ2luZ2ZhY2Uva2VybmVscy1jb21tdW5pdHk6cmVmOnJlZnMvaGVhZHMvbWFpbjCBigYKKwYBBAHWeQIEAgR8BHoAeAB2AN09MGrGxxEyYxkeHJlnNwKiSl643jyt/4eKcoAvKe6OAAABoMTNaZ4AAAQDAEcwRQIgP5gxSfLy4FLP9vZ4T4C2HBN01Vp/ZNNmXDiNJ3WZD6ACIQCMqPG8YUNpvBoo+si9YGfqD0VwIvi7PJSA1P2G8UfREDAKBggqhkjOPQQDAwNoADBlAjEA1D95XvsUGbnqjYBmuXYNdygeRfJhWSWI3AhvR46OC/CYXaqPHkFguedK8W5xUYhIAjAHWbzokgYBbsMSWbI621lxdeUfGGh4BKLr4z6Ka17axaCjQqK1dha5RQusUp4bzL4="}, "tlogEntries":[{"logIndex":"2906397162", "logId":{"keyId":"wNI9atQGlz+VWfO6LRygH4QUfY/8W4RFwiT5i5WRgB0="}, "kindVersion":{"kind":"hashedrekord", "version":"0.0.1"}, "integratedTime":"1790008191", "inclusionPromise":{"signedEntryTimestamp":"MEQCICAibnJuErv4oTr2p9bLcj0cUuI4sYBjdKMeUNBmjEdMAiB956LiLQAgS3owSEq8A2dIW5RobiBhPv3UJjU94ulNWg=="}, "inclusionProof":{"logIndex":"2784492900", "rootHash":"ruuRuA2bOwvUCTUwjql5lLPMMxBRbVBujq7C52EY+xI=", "treeSize":"2784492912", "hashes":["W4aHnZgxOxaEX8nN0MvEbfXKv21ls3qgPIO9h5ANlQw=", "YPmXuUxpSz1W8ielAHHAKbM60+hu3794xNkvr1tqwd0=", "d2ml3/jgADqEi8yzSiS9BGj6oH8Ud81IHkRiyTHT9A0=", "dQijL74kFpEfFyGTF9IGOMoU04sZsc791MgeiLDR/ng=", "a2ewThrYCNlBctYvOMWwas2uP6zmyXk9g6mMYpGjqRA=", "vgq7Vwl4M42Xa6zosYvUzr3OX6U14I8aMd2kpNxaRAs=", "KGl94H7+f+B/yxu/2qgaWzTvKr9tDUarFyMOCuFM25c=", "n5rW0Uy02cMAl76x6jWiwI7hKyUcCC7Fr0zU6EJwOdA=", "a9uxtQgboWs7ecf47zfjx0t0FQE/PqTjpq+xX55YJSo=", "uMrI4K593Uc4YIUn568kQQU3P5WI4DcAn+rS3jG0Loo=", "NxuiinDjnbF2IyjYZfiDnZUwafSidtG8YO4Q8hE7VLk=", "0NTU/2u7wVmA0DVpAJ84n4sVvSCho8WHC+DkEdg92JU=", "zkeHshHFXmbLQPJRTS95jeMY35dwTUsGQe3ERLPQ7RE=", "XU+M8SHgYBPE6zkZMa/88OBoHCVM5MhtjfLPXnLvW7g=", "jsRQr2n8GzEgko08H6qJKNeX3a7bkPYvDuB55qbMm7g=", "4N5aDWlTGjrYD2P9BkugCVYpeek+qqUf//YGGQZkIsE=", "zA+3WHfvsHMHgoq3paLzdllP+3iYxHkVPro8CUBuZa4=", "GfSJjUNAa3XTMUgDFMrv4pTzWdeXIUdphT8VAloU9uM=", "/2a58lzrUQfoAuFNUx07Dmjo61/cNyxjKuxmZfV/uSU=", "tyO4Oj1UrXex15tzrEEO6sSnE8y8obZWghZ2XixBPtE=", "rjDiNJHGnboMYYUH7OSezbKWj3wMXRTHGCLa/SaZ6O4=", "mfWjSdNaEfNXxDww6AAOw/NoK9e6z1yKaZm5xvmMabs=", "qxzHanAzz57SDdmJe0B7bJK72NTIBbwEMGBBKvDOROw=", "xH/DCseLHr9eKoYT8qsORZK7zVdEGYWHuVtsVrD95wY="], "checkpoint":{"envelope":"rekor.sigstore.dev - 1193050959916656506\n2784492912\nruuRuA2bOwvUCTUwjql5lLPMMxBRbVBujq7C52EY+xI=\n\n— rekor.sigstore.dev wNI9ajBFAiAEwubxDVjUNcv+3VF0UEaDoWvSKfocgRHNTlIm7SbBIAIhAK5uuPxv1Jv8MIC9e5jdi/yn2wE+dcljjKFJA2yCOvEc\n"}}, "canonicalizedBody":"eyJhcGlWZXJzaW9uIjoiMC4wLjEiLCJraW5kIjoiaGFzaGVkcmVrb3JkIiwic3BlYyI6eyJkYXRhIjp7Imhhc2giOnsiYWxnb3JpdGhtIjoic2hhMjU2IiwidmFsdWUiOiI2YmM4NDE1NDE2MTYyN2MyMTJlYjg5YmM5MTMzMTI5NDVjYmRkMDQxMDAyYWFhY2I5YWZkNGMxMmRkODFlNzRhIn19LCJzaWduYXR1cmUiOnsiY29udGVudCI6Ik1FVUNJQ0VJOE0vdlVKZEJVTXNUbTVoU2p1TGEvTVNJNlRkdGpZdkp2c0pWYjVvSEFpRUFnaTFGdGNnRDFpQm5oQjVtWEI2Z1g1VkVvK2hqbnpCam8xQW5iY3k1MWtnPSIsInB1YmxpY0tleSI6eyJjb250ZW50IjoiTFMwdExTMUNSVWRKVGlCRFJWSlVTVVpKUTBGVVJTMHRMUzB0Q2sxSlNVaFRla05EUW5SSFowRjNTVUpCWjBsVldqZ3hjMjU2YW05SWJrcDVlR2RLTkhBM1FUbDVWM1I0UldJMGQwTm5XVWxMYjFwSmVtb3dSVUYzVFhjS1RucEZWazFDVFVkQk1WVkZRMmhOVFdNeWJHNWpNMUoyWTIxVmRWcEhWakpOVWpSM1NFRlpSRlpSVVVSRmVGWjZZVmRrZW1SSE9YbGFVekZ3WW01U2JBcGpiVEZzV2tkc2FHUkhWWGRJYUdOT1RXcFpkMDlVU1hoTlZGbDVUMVJWZUZkb1kwNU5hbGwzVDFSSmVFMVVXWHBQVkZWNFYycEJRVTFHYTNkRmQxbElDa3R2V2tsNmFqQkRRVkZaU1V0dldrbDZhakJFUVZGalJGRm5RVVUzVnpWVVRIQktaMVprUWsweVJpOU9SbWhPYUVOblNUTkxRMlF4TWxobFdIUlJNbklLZGpsMmMzQnRVME42V2pSSFVVdzBaV1JRVG1ZNGJYWnVXbkJPTWtGYWRteDZjbXgxYWpOM2JWaExXRUl2WW5CNVRFdFBRMEptUVhkbloxaHpUVUUwUndwQk1WVmtSSGRGUWk5M1VVVkJkMGxJWjBSQlZFSm5UbFpJVTFWRlJFUkJTMEpuWjNKQ1owVkdRbEZqUkVGNlFXUkNaMDVXU0ZFMFJVWm5VVlVyZVhkT0NtbGlPSEpRWlU4eFFXNUdWVGRDY3pKMmEzZG1kVzlCZDBoM1dVUldVakJxUWtKbmQwWnZRVlV6T1ZCd2VqRlphMFZhWWpWeFRtcHdTMFpYYVhocE5Ga0tXa1E0ZDJGM1dVUldVakJTUVZGSUwwSkhSWGRZTkZwa1lVaFNNR05JVFRaTWVUbHVZVmhTYjJSWFNYVlpNamwwVERKb01Wb3laSEJpYldSdFdWZE9iQXBNTW5Sc1kyMDFiR0pJVFhSWk1qbDBZbGhXZFdGWVVqVk1lVFZ1WVZoU2IyUlhTWFprTWpsNVlUSmFjMkl6WkhwTU1rb3hZVmQ0YTB4dWJHaGlWM2hCQ21OdFZtMWplVGx2V2xkR2EyTjVPWFJaVjJ4MVRVUnJSME5wYzBkQlVWRkNaemM0ZDBGUlJVVkxNbWd3WkVoQ2VrOXBPSFprUnpseVdsYzBkVmxYVGpBS1lWYzVkV041Tlc1aFdGSnZaRmRLTVdNeVZubFpNamwxWkVkV2RXUkROV3BpTWpCM1NIZFpTMHQzV1VKQ1FVZEVkbnBCUWtGblVWSmtNamw1WVRKYWN3cGlNMlJtV2tkc2VtTkhSakJaTW1kM1RtZFpTMHQzV1VKQ1FVZEVkbnBCUWtGM1VXOVpha0pxVG1wQk5FMUVSbXhPVkZFMVRrUnNhazlVU1hwT1JGWnFDazFIUlRST2JVVjZUVzFaZWs1dFJYaGFiVkY1VDBSbk1WcHFRVlJDWjI5eVFtZEZSVUZaVHk5TlFVVkZRa0ZXUTJSWGJITmFSRUZ5UW1kdmNrSm5SVVVLUVZsUEwwMUJSVVpDUWpGdlpGZGtibUZYTlc1YWJVWnFXbE01Y2xwWVNuVmFWM2g2VEZkT2RtSlhNVEZpYld3d1pWUkJaRUpuYjNKQ1owVkZRVmxQTHdwTlFVVkhRa0U1ZVZwWFducE1NbWhzV1ZkU2Vrd3lNV2hoVnpSM1QzZFpTMHQzV1VKQ1FVZEVkbnBCUWtOQlVYUkVRM1J2WkVoU2QyTjZiM1pNTTFKMkNtRXlWblZNYlVacVpFZHNkbUp1VFhWYU1td3dZVWhXYVdSWVRteGpiVTUyWW01U2JHSnVVWFZaTWpsMFRVY3dSME5wYzBkQlVWRkNaemM0ZDBGUmEwVUtXSGQ0WkdGSVVqQmpTRTAyVEhrNWJtRllVbTlrVjBsMVdUSTVkRXd5YURGYU1tUndZbTFrYlZsWFRteE1NblJzWTIwMWJHSklUWFJaTWpsMFlsaFdkUXBoV0ZJMVRIazFibUZZVW05a1YwbDJaREk1ZVdFeVduTmlNMlI2VERKS01XRlhlR3RNYm14b1lsZDRRV050Vm0xamVUbHZXbGRHYTJONU9YUlpWMngxQ2sxRVowZERhWE5IUVZGUlFtYzNPSGRCVVc5RlMyZDNiMWxxUW1wT2FrRTBUVVJHYkU1VVVUVk9SR3hxVDFSSmVrNUVWbXBOUjBVMFRtMUZlazF0V1hvS1RtMUZlRnB0VVhsUFJHY3hXbXBCWWtKbmIzSkNaMFZGUVZsUEwwMUJSVXhDUVRCTlF6Tk9iR0pIV1hSaFJ6bDZaRWRXYTAxRlFVZERhWE5IUVZGUlFncG5OemgzUVZGM1JVMW5kM2RoU0ZJd1kwaE5Oa3g1T1c1aFdGSnZaRmRKZFZreU9YUk1NbWd4V2pKa2NHSnRaRzFaVjA1c1RESjBiR050Tld4aVNFMTBDbGt5T1hSaVdGWjFZVmhTTlUxRVowZERhWE5IUVZGUlFtYzNPSGRCVVRCRlMyZDNiMWxxUW1wT2FrRTBUVVJHYkU1VVVUVk9SR3hxVDFSSmVrNUVWbW9LVFVkRk5FNXRSWHBOYlZsNlRtMUZlRnB0VVhsUFJHY3hXbXBCWmtKbmIzSkNaMFZGUVZsUEwwMUJSVTlDUWtWTlJETktiRnB1VFhaaFIxWm9Xa2hOZGdwaVYwWndZbXBCWVVKbmIzSkNaMFZGUVZsUEwwMUJSVkJDUVhkTlEycEZkMDU2UlRCT2VsVXhUV3ByZDB4bldVdExkMWxDUWtGSFJIWjZRVUpGUVZGbkNrUkNOVzlrU0ZKM1kzcHZka3d5WkhCa1IyZ3hXV2sxYW1JeU1IWmhTRlp1V2pKc2RWb3lXbWhaTWxWM1IwRlpTMHQzV1VKQ1FVZEVkbnBCUWtWUlVVc0tSRUZuZVU1VVkzbE5SR013VFhwQ2RFSm5iM0pDWjBWRlFWbFBMMDFCUlZOQ1JqaE5XRmRvTUdSSVFucFBhVGgyV2pKc01HRklWbWxNYlU1MllsTTVid3BrVjJSdVlWYzFibHB0Um1wYVV6bHlXbGhLZFZwWGVIcE1WMDUyWWxjeE1XSnRiREJsVXpoMVdqSnNNR0ZJVm1sTU0yUjJZMjEwYldKSE9UTmplVGxwQ21SWGJITmFRelUxV1ZjeGMxRklTbXhhYmsxMllVZFdhRnBJVFhaaVYwWndZbXBCTkVKbmIzSkNaMFZGUVZsUEwwMUJSVlJDUTI5TlMwZEpkMWw2V1hjS1QwUkJlRnBVVlRCUFZGRTFXWHByZVUxNlVURlpla0pvVDBSYWFFMTZTbTFOZWxwb1RWZGFhMDFxWnpST1YxbDNTVkZaUzB0M1dVSkNRVWRFZG5wQlFncEdRVkZVUkVKR00ySXpTbkphYlhoMlpERTVhMkZZVG5kWldGSnFZVVJDYTBKbmIzSkNaMFZGUVZsUEwwMUJSVlpDUmxsTlZrZG9NR1JJUW5wUGFUaDJDbG95YkRCaFNGWnBURzFPZG1KVE9XOWtWMlJ1WVZjMWJscHRSbXBhVXpseVdsaEtkVnBYZUhwTVYwNTJZbGN4TVdKdGJEQmxVemxvV1ROU2NHSXlOWG9LVEROS01XSnVUWFpOZWxVeVRXcFZlazFxVFRWTlZHTjJXVmhTTUZwWE1YZGtTRTEyVFZSQlYwSm5iM0pDWjBWRlFWbFBMMDFCUlZkQ1FXZE5RbTVDTVFwWmJYaHdXWHBDUjBKbmIzSkNaMFZGUVZsUEwwMUJSVmxDUkdkTlRtNUtiR05IT0RaaFNGWnVXakpzZFZveVdtaFpNbFYyWVRKV2VXSnRWbk5qZVRGcUNtSXlNWFJrVnpWd1pFaHJObU50Vm0xUGJrcHNXbTVOZG1GSFZtaGFTRTEyWWxkR2NHSnFRMEpwWjFsTFMzZFpRa0pCU0ZkbFVVbEZRV2RTT0VKSWIwRUtaVUZDTWtGT01EbE5SM0pIZUhoRmVWbDRhMlZJU214dVRuZExhVk5zTmpRemFubDBMelJsUzJOdlFYWkxaVFpQUVVGQlFtOU5WRTVoV2pSQlFVRlJSQXBCUldOM1VsRkpaMUExWjNoVFpreDVORVpNVURsMldqUlVORU15U0VKT01ERldjQzlhVGs1dFdFUnBUa296VjFwRU5rRkRTVkZEVFhGUVJ6aFpWVTV3Q25aQ2IyOHJjMms1V1VkbWNVUXdWbmRKZG1rM1VFcFRRVEZRTWtjNFZXWlNSVVJCUzBKblozRm9hMnBQVUZGUlJFRjNUbTlCUkVKc1FXcEZRVEZFT1RVS1dIWnpWVWRpYm5GcVdVSnRkVmhaVG1SNVoyVlNaa3BvVjFOWFNUTkJhSFpTTkRaUFF5OURXVmhoY1ZCSWEwWm5kV1ZrU3poWE5YaFZXV2hKUVdwQlNBcFhZbnB2YTJkWlFtSnpUVk5YWWtrMk1qRnNlR1JsVldaSFIyZzBRa3RNY2pSNk5rdGhNVGRoZUdGRGFsRnhTekZrYUdFMVVsRjFjMVZ3TkdKNlREUTlDaTB0TFMwdFJVNUVJRU5GVWxSSlJrbERRVlJGTFMwdExTMEsifX19fQ=="}], "timestampVerificationData":{"rfc3161Timestamps":[{"signedTimestamp":"MIICyjADAgEAMIICwQYJKoZIhvcNAQcCoIICsjCCAq4CAQMxDTALBglghkgBZQMEAgEwgbgGCyqGSIb3DQEJEAEEoIGoBIGlMIGiAgEBBgkrBgEEAYO/MAIwMTANBglghkgBZQMEAgEFAAQgDM0ZH/BFkL7abcswmvCpzKpHqo/dDY4oQpTE2rSegcYCFQCEQUOszsswNLq5HwXLmDQPi0uCoBgPMjAyNjA5MjExNjI5NTFaMAMCAQGgMqQwMC4xFTATBgNVBAoTDHNpZ3N0b3JlLmRldjEVMBMGA1UEAxMMc2lnc3RvcmUtdHNhoAAxggHbMIIB1wIBATBRMDkxFTATBgNVBAoTDHNpZ3N0b3JlLmRldjEgMB4GA1UEAxMXc2lnc3RvcmUtdHNhLXNlbGZzaWduZWQCFDoTVC8MkGHuvMFDL8uKjosqI4sMMAsGCWCGSAFlAwQCAaCB/DAaBgkqhkiG9w0BCQMxDQYLKoZIhvcNAQkQAQQwHAYJKoZIhvcNAQkFMQ8XDTI2MDkyMTE2Mjk1MVowLwYJKoZIhvcNAQkEMSIEINZVjoB6Ftn+joLxXvUK/3YRG/5j5NvBF1Bn3jv91UpoMIGOBgsqhkiG9w0BCRACLzF/MH0wezB5BCCF+Se8B6tiysO0Q1bBDvyBssaIP9p6uebYcNnROs0FtzBVMD2kOzA5MRUwEwYDVQQKEwxzaWdzdG9yZS5kZXYxIDAeBgNVBAMTF3NpZ3N0b3JlLXRzYS1zZWxmc2lnbmVkAhQ6E1QvDJBh7rzBQy/Lio6LKiOLDDAKBggqhkjOPQQDAgRnMGUCMHuLT4V9Ga1MtwJ9bd3spYCretPRwz5W9+Qb9Vn2yXVXqo5BcKaoV9k9a6ad7UZwogIxAI8j8+evxPw2yj9+byUuefD5+AGCAppHEIPdzwM0Gn04tiTfiWKL4FeqQXga/7aB4w=="}]}}, "messageSignature":{"messageDigest":{"algorithm":"SHA2_256", "digest":"a8hBVBYWJ8IS64m8kTMSlFy90EEAKqrLmv1MEt2B50o="}, "signature":"MEUCICEI8M/vUJdBUMsTm5hSjuLa/MSI6TdtjYvJvsJVb5oHAiEAgi1FtcgD1iBnhB5mXB6gX5VEo+hjnzBjo1Anbcy51kg="}}
build/torch-rocm/moe.py CHANGED
@@ -47,7 +47,15 @@ from triton.language.extra.cuda import gdc_launch_dependents, gdc_wait
47
  from .grouped import matmul_grouped
48
  from .batched import GATE_UNSTACK_MAX_S, matmul_batched
49
  from .bayesian_autotuner import bayesian_autotune
50
- from .compat import MX_SCALE_GROUP_K, NVFP4_SCALE_GROUP_K, compile_time_only_triton_wrap, decode_pdl, device_context
 
 
 
 
 
 
 
 
51
  from .formats import get_supported_act_fns, is_mx, is_mxfp4, weight_format
52
  from .norm import norm_column_factor, rms_inv_rows, rms_norm_rows
53
  from .quant import _launch_act_quant
@@ -320,6 +328,8 @@ def moe_fused_grouped(
320
  down_proj_weight_global_scale: torch.Tensor | None = None,
321
  gate_up_proj_input_global_scale: torch.Tensor | None = None,
322
  down_proj_input_global_scale: torch.Tensor | None = None,
 
 
323
  post_expert_norm=None, # the model's per-expert output norm on the routed rows: a
324
  # get_supported_norms() name (fused into the reduce) or a host callable
325
  post_expert_norm_weight: torch.Tensor | None = None, # (H,) weight of a named norm
@@ -363,9 +373,11 @@ def moe_fused_grouped(
363
  # intermediate (the op quantizes the raw hidden itself and owns the expand-vs-gather
364
  # regime policy — this forward is pure sequencing). scatter_idx=None: the down reads
365
  # the intermediate in place. (C, Cs) under a requant format; a bare Tensor otherwise.
 
366
  gate_up_out = matmul_grouped(
367
  hidden_states,
368
  gate_up_proj,
 
369
  Bs=gate_up_proj_scale_inv,
370
  a_global_scale=gate_up_proj_input_global_scale,
371
  b_global_scale=gate_up_proj_weight_global_scale,
@@ -376,9 +388,9 @@ def moe_fused_grouped(
376
  **glu,
377
  bias=gate_up_proj_bias,
378
  # fmt is the resolved format; "bf16" (weight-only) leaves the GLU intermediate bf16,
379
- # no requant.
380
  activation_format=fmt,
381
- quantize_output=bool(glu) and fmt != "bf16",
382
  output_dtype=hidden_states.dtype,
383
  gather_idx=gather_idx,
384
  )
@@ -392,7 +404,7 @@ def moe_fused_grouped(
392
  down_out = matmul_grouped(
393
  inter,
394
  down_proj,
395
- As=inter_scale,
396
  Bs=down_proj_scale_inv,
397
  a_global_scale=down_proj_input_global_scale,
398
  b_global_scale=down_proj_weight_global_scale,
@@ -429,6 +441,8 @@ def moe_fused_batched(
429
  down_proj_weight_global_scale: torch.Tensor | None = None,
430
  gate_up_proj_input_global_scale: torch.Tensor | None = None,
431
  down_proj_input_global_scale: torch.Tensor | None = None,
 
 
432
  post_expert_norm=None, # the model's per-expert output norm on the routed rows: a
433
  # get_supported_norms() name (fused into the reduce) or a host callable
434
  post_expert_norm_weight: torch.Tensor | None = None, # (H,) weight of a named norm
@@ -472,9 +486,11 @@ def moe_fused_batched(
472
  # intermediate (the op quantizes the raw activations). gather_idx reads each routed
473
  # row from the unexpanded hidden in-kernel (no copy).
474
  # (C, Cs) under a requant format; a bare Tensor on the full-precision path
 
475
  gate_up_out = matmul_batched(
476
  hidden_states,
477
  gate_up_proj,
 
478
  Bs=gate_up_proj_scale_inv,
479
  a_global_scale=gate_up_proj_input_global_scale,
480
  b_global_scale=gate_up_proj_weight_global_scale,
@@ -488,10 +504,14 @@ def moe_fused_batched(
488
  # (``fused_glu(quant_group=...)`` — one launch, hands the down a ready fp8+scales intermediate and
489
  # kills its offline act quant); ABOVE the band the stacked epilogue's requant pins
490
  # the gate|up tile to the whole block scale and halves the grid, so the bf16 handoff
491
- # (down inline-quants) stays the win there.
 
492
  activation_format=fmt,
493
  quantize_output=(
494
- bool(glu) and fmt != "bf16" and (fmt != "fp8" or expert_ids.numel() <= GATE_UNSTACK_MAX_S)
 
 
 
495
  ),
496
  output_dtype=hidden_states.dtype,
497
  gather_idx=gather_idx,
@@ -506,7 +526,7 @@ def moe_fused_batched(
506
  down_out = matmul_batched(
507
  inter,
508
  down_proj,
509
- As=inter_scale,
510
  Bs=down_proj_scale_inv,
511
  a_global_scale=down_proj_input_global_scale,
512
  b_global_scale=down_proj_weight_global_scale,
@@ -545,6 +565,8 @@ def moe_unfused_grouped(
545
  down_proj_weight_global_scale: torch.Tensor | None = None,
546
  gate_up_proj_input_global_scale: torch.Tensor | None = None,
547
  down_proj_input_global_scale: torch.Tensor | None = None,
 
 
548
  post_expert_norm=None, # the model's per-expert output norm on the routed rows: a
549
  # get_supported_norms() name (fused into the reduce) or a host callable
550
  post_expert_norm_weight: torch.Tensor | None = None, # (H,) weight of a named norm
@@ -562,8 +584,10 @@ def moe_unfused_grouped(
562
  ``moe_fused_grouped`` but the SwiGLU + intermediate quant happen between two plain GEMMs
563
  rather than inside the gate_up epilogue; each GEMM quantizes its raw input in
564
  ``activation_format`` (``None`` follows the weight format, mirroring the fused forward — mxfp4
565
- weights run the all-fp4 W4A4 chain). All formats route through the shared ``matmul_grouped``. The NVFP4 activation globals thread the same way as the fused sibling:
566
- each GEMM quantizes its raw input against its ``*_input_global_scale``. ``act_fn`` is a
 
 
567
  ``get_supported_act_fns()`` name (fused into the gate_up epilogue where the forward fuses) or any
568
  callable applied on the host to the raw gate_up output; ``gate=False`` runs an ungated
569
  projection. Scales are affine or pre-swizzled (``SWIZZLE_32_4_4``, self-describing 5-D)
@@ -582,6 +606,7 @@ def moe_unfused_grouped(
582
  gate_up_out = matmul_grouped(
583
  hidden_states,
584
  gate_up_proj,
 
585
  Bs=gate_up_proj_scale_inv,
586
  a_global_scale=gate_up_proj_input_global_scale,
587
  b_global_scale=gate_up_proj_weight_global_scale,
@@ -597,6 +622,7 @@ def moe_unfused_grouped(
597
  down_out = matmul_grouped(
598
  inter,
599
  down_proj,
 
600
  Bs=down_proj_scale_inv,
601
  a_global_scale=down_proj_input_global_scale,
602
  b_global_scale=down_proj_weight_global_scale,
@@ -624,6 +650,8 @@ def moe_torch_grouped(
624
  down_proj_weight_global_scale: torch.Tensor | None = None,
625
  gate_up_proj_input_global_scale: torch.Tensor | None = None,
626
  down_proj_input_global_scale: torch.Tensor | None = None,
 
 
627
  post_expert_norm=None, # the model's per-expert output norm on the routed rows: a
628
  # get_supported_norms() name (fused into the reduce) or a host callable
629
  post_expert_norm_weight: torch.Tensor | None = None, # (H,) weight of a named norm
@@ -653,6 +681,10 @@ def moe_torch_grouped(
653
  ACTIVATION scale (which changes each call). The format is read off the dtypes (the block
654
  preserves them: E4M3 scale = NVFP4, uint8 = MX; packed-E2M1 weight = int8) since the blocked
655
  shape no longer matches the group-shape detectors."""
 
 
 
 
656
  assert gate_up_proj.dtype in (torch.int8, torch.float8_e4m3fn), (
657
  "torch grouped baseline is MX-only (packed E2M1 or E4M3 weights)"
658
  )
@@ -664,9 +696,10 @@ def moe_torch_grouped(
664
  "the torch baseline always quantizes activations (scaled_grouped_mm has no bf16-act x "
665
  "MX-weight form) — activation_format='bf16' (W4A16/W8A16) is not representable here"
666
  )
667
-
668
- import torch.nn.functional as F
669
- from torch.nn.functional import ScalingType, SwizzleType
 
670
 
671
  # torchao >= 0.18 required with cutlass-dsl >= 4.6 (0.17 imports a helper path 4.6
672
  # removed; fixed upstream in pytorch/ao).
@@ -803,6 +836,8 @@ def moe_unfused_batched(
803
  down_proj_weight_global_scale: torch.Tensor | None = None,
804
  gate_up_proj_input_global_scale: torch.Tensor | None = None,
805
  down_proj_input_global_scale: torch.Tensor | None = None,
 
 
806
  post_expert_norm=None, # the model's per-expert output norm on the routed rows: a
807
  # get_supported_norms() name (fused into the reduce) or a host callable
808
  post_expert_norm_weight: torch.Tensor | None = None, # (H,) weight of a named norm
@@ -819,8 +854,10 @@ def moe_unfused_batched(
819
  down (plain batched GEMM) → routing-weighted reduce. Same math as ``moe_fused_batched`` but
820
  the SwiGLU + intermediate quant happen between two plain GEMMs; each GEMM quantizes its raw
821
  input in ``activation_format`` (``None`` follows the weight format, ``"bf16"`` is weight-only). All
822
- formats route through the shared ``matmul_batched``. The NVFP4 activation globals thread the same way as the fused sibling:
823
- each GEMM quantizes its raw input against its ``*_input_global_scale``. ``act_fn`` is a
 
 
824
  ``get_supported_act_fns()`` name (fused into the gate_up epilogue where the forward fuses) or any
825
  callable applied on the host to the raw gate_up output; ``gate=False`` runs an ungated
826
  projection. Scales are affine or pre-swizzled (``SWIZZLE_32_4_4``, self-describing 5-D)
@@ -836,6 +873,7 @@ def moe_unfused_batched(
836
  gate_up_out = matmul_batched(
837
  hidden_states,
838
  gate_up_proj,
 
839
  Bs=gate_up_proj_scale_inv,
840
  a_global_scale=gate_up_proj_input_global_scale,
841
  b_global_scale=gate_up_proj_weight_global_scale,
@@ -850,6 +888,7 @@ def moe_unfused_batched(
850
  down_out = matmul_batched(
851
  inter,
852
  down_proj,
 
853
  Bs=down_proj_scale_inv,
854
  a_global_scale=down_proj_input_global_scale,
855
  b_global_scale=down_proj_weight_global_scale,
 
47
  from .grouped import matmul_grouped
48
  from .batched import GATE_UNSTACK_MAX_S, matmul_batched
49
  from .bayesian_autotuner import bayesian_autotune
50
+ from .compat import (
51
+ MX_SCALE_GROUP_K,
52
+ NVFP4_SCALE_GROUP_K,
53
+ ScalingType,
54
+ SwizzleType,
55
+ compile_time_only_triton_wrap,
56
+ decode_pdl,
57
+ device_context,
58
+ )
59
  from .formats import get_supported_act_fns, is_mx, is_mxfp4, weight_format
60
  from .norm import norm_column_factor, rms_inv_rows, rms_norm_rows
61
  from .quant import _launch_act_quant
 
328
  down_proj_weight_global_scale: torch.Tensor | None = None,
329
  gate_up_proj_input_global_scale: torch.Tensor | None = None,
330
  down_proj_input_global_scale: torch.Tensor | None = None,
331
+ gate_up_proj_activation_scale: torch.Tensor | None = None,
332
+ down_proj_activation_scale: torch.Tensor | None = None,
333
  post_expert_norm=None, # the model's per-expert output norm on the routed rows: a
334
  # get_supported_norms() name (fused into the reduce) or a host callable
335
  post_expert_norm_weight: torch.Tensor | None = None, # (H,) weight of a named norm
 
373
  # intermediate (the op quantizes the raw hidden itself and owns the expand-vs-gather
374
  # regime policy — this forward is pure sequencing). scatter_idx=None: the down reads
375
  # the intermediate in place. (C, Cs) under a requant format; a bare Tensor otherwise.
376
+ static_act = gate_up_proj_activation_scale is not None or down_proj_activation_scale is not None
377
  gate_up_out = matmul_grouped(
378
  hidden_states,
379
  gate_up_proj,
380
+ As=gate_up_proj_activation_scale,
381
  Bs=gate_up_proj_scale_inv,
382
  a_global_scale=gate_up_proj_input_global_scale,
383
  b_global_scale=gate_up_proj_weight_global_scale,
 
388
  **glu,
389
  bias=gate_up_proj_bias,
390
  # fmt is the resolved format; "bf16" (weight-only) leaves the GLU intermediate bf16,
391
+ # no requant. So does static, whose epilogue has no calibrated scale to requant against.
392
  activation_format=fmt,
393
+ quantize_output=bool(glu) and fmt != "bf16" and not static_act,
394
  output_dtype=hidden_states.dtype,
395
  gather_idx=gather_idx,
396
  )
 
404
  down_out = matmul_grouped(
405
  inter,
406
  down_proj,
407
+ As=down_proj_activation_scale if static_act else inter_scale,
408
  Bs=down_proj_scale_inv,
409
  a_global_scale=down_proj_input_global_scale,
410
  b_global_scale=down_proj_weight_global_scale,
 
441
  down_proj_weight_global_scale: torch.Tensor | None = None,
442
  gate_up_proj_input_global_scale: torch.Tensor | None = None,
443
  down_proj_input_global_scale: torch.Tensor | None = None,
444
+ gate_up_proj_activation_scale: torch.Tensor | None = None,
445
+ down_proj_activation_scale: torch.Tensor | None = None,
446
  post_expert_norm=None, # the model's per-expert output norm on the routed rows: a
447
  # get_supported_norms() name (fused into the reduce) or a host callable
448
  post_expert_norm_weight: torch.Tensor | None = None, # (H,) weight of a named norm
 
486
  # intermediate (the op quantizes the raw activations). gather_idx reads each routed
487
  # row from the unexpanded hidden in-kernel (no copy).
488
  # (C, Cs) under a requant format; a bare Tensor on the full-precision path
489
+ static_act = gate_up_proj_activation_scale is not None or down_proj_activation_scale is not None
490
  gate_up_out = matmul_batched(
491
  hidden_states,
492
  gate_up_proj,
493
+ As=gate_up_proj_activation_scale,
494
  Bs=gate_up_proj_scale_inv,
495
  a_global_scale=gate_up_proj_input_global_scale,
496
  b_global_scale=gate_up_proj_weight_global_scale,
 
504
  # (``fused_glu(quant_group=...)`` — one launch, hands the down a ready fp8+scales intermediate and
505
  # kills its offline act quant); ABOVE the band the stacked epilogue's requant pins
506
  # the gate|up tile to the whole block scale and halves the grid, so the bf16 handoff
507
+ # (down inline-quants) stays the win there. Static keeps bf16 either way, its epilogue
508
+ # having no calibrated scale to requant against.
509
  activation_format=fmt,
510
  quantize_output=(
511
+ bool(glu)
512
+ and fmt != "bf16"
513
+ and (fmt != "fp8" or expert_ids.numel() <= GATE_UNSTACK_MAX_S)
514
+ and not static_act
515
  ),
516
  output_dtype=hidden_states.dtype,
517
  gather_idx=gather_idx,
 
526
  down_out = matmul_batched(
527
  inter,
528
  down_proj,
529
+ As=down_proj_activation_scale if static_act else inter_scale,
530
  Bs=down_proj_scale_inv,
531
  a_global_scale=down_proj_input_global_scale,
532
  b_global_scale=down_proj_weight_global_scale,
 
565
  down_proj_weight_global_scale: torch.Tensor | None = None,
566
  gate_up_proj_input_global_scale: torch.Tensor | None = None,
567
  down_proj_input_global_scale: torch.Tensor | None = None,
568
+ gate_up_proj_activation_scale: torch.Tensor | None = None,
569
+ down_proj_activation_scale: torch.Tensor | None = None,
570
  post_expert_norm=None, # the model's per-expert output norm on the routed rows: a
571
  # get_supported_norms() name (fused into the reduce) or a host callable
572
  post_expert_norm_weight: torch.Tensor | None = None, # (H,) weight of a named norm
 
584
  ``moe_fused_grouped`` but the SwiGLU + intermediate quant happen between two plain GEMMs
585
  rather than inside the gate_up epilogue; each GEMM quantizes its raw input in
586
  ``activation_format`` (``None`` follows the weight format, mirroring the fused forward — mxfp4
587
+ weights run the all-fp4 W4A4 chain). All formats route through the shared ``matmul_grouped``.
588
+ The NVFP4 activation globals thread the same way as the fused sibling: each GEMM quantizes its
589
+ raw input against its ``*_input_global_scale``, and a ``*_activation_scale`` quantizes it
590
+ against that calibrated scale instead of a runtime one. ``act_fn`` is a
591
  ``get_supported_act_fns()`` name (fused into the gate_up epilogue where the forward fuses) or any
592
  callable applied on the host to the raw gate_up output; ``gate=False`` runs an ungated
593
  projection. Scales are affine or pre-swizzled (``SWIZZLE_32_4_4``, self-describing 5-D)
 
606
  gate_up_out = matmul_grouped(
607
  hidden_states,
608
  gate_up_proj,
609
+ As=gate_up_proj_activation_scale,
610
  Bs=gate_up_proj_scale_inv,
611
  a_global_scale=gate_up_proj_input_global_scale,
612
  b_global_scale=gate_up_proj_weight_global_scale,
 
622
  down_out = matmul_grouped(
623
  inter,
624
  down_proj,
625
+ As=down_proj_activation_scale,
626
  Bs=down_proj_scale_inv,
627
  a_global_scale=down_proj_input_global_scale,
628
  b_global_scale=down_proj_weight_global_scale,
 
650
  down_proj_weight_global_scale: torch.Tensor | None = None,
651
  gate_up_proj_input_global_scale: torch.Tensor | None = None,
652
  down_proj_input_global_scale: torch.Tensor | None = None,
653
+ gate_up_proj_activation_scale: torch.Tensor | None = None,
654
+ down_proj_activation_scale: torch.Tensor | None = None,
655
  post_expert_norm=None, # the model's per-expert output norm on the routed rows: a
656
  # get_supported_norms() name (fused into the reduce) or a host callable
657
  post_expert_norm_weight: torch.Tensor | None = None, # (H,) weight of a named norm
 
681
  ACTIVATION scale (which changes each call). The format is read off the dtypes (the block
682
  preserves them: E4M3 scale = NVFP4, uint8 = MX; packed-E2M1 weight = int8) since the blocked
683
  shape no longer matches the group-shape detectors."""
684
+ assert ScalingType is not None, (
685
+ "this torch has no torch.nn.functional.ScalingType — the baseline's scaled_grouped_mm "
686
+ "scaling enums arrived with the op itself, so an older torch cannot run it"
687
+ )
688
  assert gate_up_proj.dtype in (torch.int8, torch.float8_e4m3fn), (
689
  "torch grouped baseline is MX-only (packed E2M1 or E4M3 weights)"
690
  )
 
696
  "the torch baseline always quantizes activations (scaled_grouped_mm has no bf16-act x "
697
  "MX-weight form) — activation_format='bf16' (W4A16/W8A16) is not representable here"
698
  )
699
+ assert gate_up_proj_activation_scale is None and down_proj_activation_scale is None, (
700
+ "the torch baseline quantizes each activation against a scale it derives per call — a "
701
+ "calibrated (static) scale is not representable here"
702
+ )
703
 
704
  # torchao >= 0.18 required with cutlass-dsl >= 4.6 (0.17 imports a helper path 4.6
705
  # removed; fixed upstream in pytorch/ao).
 
836
  down_proj_weight_global_scale: torch.Tensor | None = None,
837
  gate_up_proj_input_global_scale: torch.Tensor | None = None,
838
  down_proj_input_global_scale: torch.Tensor | None = None,
839
+ gate_up_proj_activation_scale: torch.Tensor | None = None,
840
+ down_proj_activation_scale: torch.Tensor | None = None,
841
  post_expert_norm=None, # the model's per-expert output norm on the routed rows: a
842
  # get_supported_norms() name (fused into the reduce) or a host callable
843
  post_expert_norm_weight: torch.Tensor | None = None, # (H,) weight of a named norm
 
854
  down (plain batched GEMM) → routing-weighted reduce. Same math as ``moe_fused_batched`` but
855
  the SwiGLU + intermediate quant happen between two plain GEMMs; each GEMM quantizes its raw
856
  input in ``activation_format`` (``None`` follows the weight format, ``"bf16"`` is weight-only). All
857
+ formats route through the shared ``matmul_batched``. The NVFP4 activation globals thread the
858
+ same way as the fused sibling: each GEMM quantizes its raw input against its
859
+ ``*_input_global_scale``, and a ``*_activation_scale`` quantizes it against that calibrated
860
+ scale instead of a runtime one. ``act_fn`` is a
861
  ``get_supported_act_fns()`` name (fused into the gate_up epilogue where the forward fuses) or any
862
  callable applied on the host to the raw gate_up output; ``gate=False`` runs an ungated
863
  projection. Scales are affine or pre-swizzled (``SWIZZLE_32_4_4``, self-describing 5-D)
 
873
  gate_up_out = matmul_batched(
874
  hidden_states,
875
  gate_up_proj,
876
+ As=gate_up_proj_activation_scale,
877
  Bs=gate_up_proj_scale_inv,
878
  a_global_scale=gate_up_proj_input_global_scale,
879
  b_global_scale=gate_up_proj_weight_global_scale,
 
888
  down_out = matmul_batched(
889
  inter,
890
  down_proj,
891
+ As=down_proj_activation_scale,
892
  Bs=down_proj_scale_inv,
893
  a_global_scale=down_proj_input_global_scale,
894
  b_global_scale=down_proj_weight_global_scale,
build/torch-rocm/pruners.py CHANGED
@@ -15,7 +15,7 @@
15
  import torch
16
  import triton
17
 
18
- from .compat import get_active_device_type, is_sm10x, is_sm90, sm_count, sm_shared_memory_limit
19
  from .mma import MMA_N_ATOM_WIDTH
20
 
21
  # ── config pruners ────────────────────────────────────────────────────────────
@@ -101,6 +101,10 @@ from .mma import MMA_N_ATOM_WIDTH
101
  # it fit shared memory and failed only as benign launch-time smem overflows.
102
  SM10X_SCALED_MMA_MAX_N = 256
103
 
 
 
 
 
104
  # Branch axes that RELOCATE the tile optimum (compute unit / operand orientation): the
105
  # tuner's guaranteed max-tile anchors group by these (``path_anchor_axes`` — a declaration,
106
  # so the tuner itself stays independent of configuration details). Scheduling axes
@@ -334,6 +338,10 @@ def mx_config_pruner(k_arg: str, n_arg: str | None = None, block_within_k: bool
334
  dot/scalar/swap arms column-unpack them to E4M3 (lossless) first — no arm is
335
  structurally packed-incompatible, so W4A4 needs no shape gate of its own.
336
 
 
 
 
 
337
  The shape gates above are scoped to ``dot_scaled`` — they are native scaled-MMA bug
338
  gates. The ``dot`` arm (BK structurally the UE8M0 group, 32) is CORRECT everywhere
339
  probed (forced-config sweep 2026-07-14, GATE and plain, MXFP4/MXFP8, incl. width-512
@@ -474,6 +482,18 @@ def mx_config_pruner(k_arg: str, n_arg: str | None = None, block_within_k: bool
474
  return (2 if args.get("GATE") else 1) * config_dim(c, args, "BLOCK_SIZE_N") >= 128
475
  return config_dim(c, args, "BLOCK_SIZE_M") >= 128
476
 
 
 
 
 
 
 
 
 
 
 
 
 
477
  def scales_are_e4m3(args):
478
  return getattr(args.get("Bs"), "dtype", None) == torch.float8_e4m3fn
479
 
@@ -495,6 +515,7 @@ def mx_config_pruner(k_arg: str, n_arg: str | None = None, block_within_k: bool
495
  # re-admitting the single-trip/width traps it existed to remove).
496
  return compose_pruners(
497
  *stages,
 
498
  config_filter(nvfp4_native_ok, when=scales_are_e4m3),
499
  config_filter(
500
  mma_trap_ok, when=lambda args: is_sm10x(), on_empty=raise_all_mma_trapped
@@ -714,49 +735,65 @@ def packed_schedule_scope_pruner(min_bm: int = 128):
714
  return config_filter(ok)
715
 
716
 
717
- def weight_only_warp_spec_matched_mode_pruner():
718
- """``early_config_prune`` for the weight-only grouped kernel's WARP_SPEC lowering wall
719
- (Triton 3.7.1 ``TritonGPUOptimizePartitionWarps`` -> PassManager::run failed; the pass
720
- cannot partition this loop when BOTH operand loads share a memory mode). Charted
721
- 2026-08-26 (48-cell forced-config matrix, BM 64 and 128 IDENTICAL — BM-independent):
722
-
723
- CM A/B modes w4 w8 w16
724
- dot ptr/ptr ok ok FAIL
725
- dot desc/desc ok ok ok (was charted FAIL/FAIL — refuted, see _TRAPPING_NUM_WARPS)
726
- dot_scaled ptr/ptr FAIL FAIL FAIL
727
- dot_scaled desc/desc ok ok ok (w8 was charted FAIL — refuted)
728
- MIXED modes (ptr/desc, desc/ptr): 24/24 ok at every warp count.
729
- BN=32 + WS: FAIL in every cell (the tile law, fenced here).
730
-
731
- The non-monotone warp dependence (w16 rescues desc/desc dot but breaks ptr/ptr dot) marks
732
- this as the pass's internal partition feasibility, not kernel source — the source is
733
- identical across passing and failing cells. Two root-cause attempts on the sibling family
734
- are refuted and recorded (block_dynamic_grouped_matmul_pruner); the kernel-side remedy is
735
- to stop emitting configs the pass cannot lower. Cost in the wild: 19 dead compiles per
736
- weight-only tune, and under inductor ONE failing config kills the whole torch.compile
737
- cell instead of scoring inf — this fence is what recovers those cells."""
738
-
739
- # desc/desc rows REMOVED 2026-08-29: forced through the tuner's own Config objects (pre_hook
740
- # intact) at the charted GPT-OSS shape, (dot, desc/desc) w4 AND w8 and (dot_scaled, desc/desc) w8
741
- # all compile and run bit-identical — dot+WS+BK=128 is the gate_up's best config (2128 vs 2235us).
742
- # The original matrix was charted with hand-built Configs, which skip the descriptor pre_hooks
743
- # and fail every descriptor cell for the wrong reason. The pointer rows stand unrefuted.
744
- _TRAPPING_NUM_WARPS = {
745
- ("dot", "pointer"): {16},
746
- ("dot_scaled", "pointer"): {4, 8, 16},
747
- }
 
 
 
 
 
 
 
 
 
 
 
 
 
748
 
749
  def ok(c, args):
750
  if not c.kwargs.get("WARP_SPEC"):
751
  return True
752
- # BN=32 + WS fails to lower in every memory mode / warp count / compute mode probed
753
- # (9/9 cells, GPT-OSS N=K=2880, 2026-08-29) — the tile law the matrix never encoded
754
- if config_dim(c, args, "BLOCK_SIZE_N") == 32:
755
  return False
756
  a, b = c.kwargs.get("A_MEMORY_MODE"), c.kwargs.get("B_MEMORY_MODE")
757
- if a != b: # mixed modes lower everywhere
 
 
 
 
758
  return True
759
- return c.num_warps not in _TRAPPING_NUM_WARPS.get((c.kwargs.get("COMPUTE_MODE"), a), set())
 
760
 
761
  return config_filter(ok, when=lambda args: is_sm10x())
762
 
@@ -776,6 +813,37 @@ def weight_only_swap_scope_pruner():
776
  return config_filter(ok)
777
 
778
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
779
  def mx_2d_swap_scope_pruner(max_m: int = 16):
780
  """``early_config_prune`` scoping the 2D mx kernel's ``SWAP_AB`` rows to the ONE regime
781
  they exist for: E4M3-scale (NVFP4) single-token decode. NVFP4 has no other native M=1
@@ -804,6 +872,14 @@ def mx_2d_swap_scope_pruner(max_m: int = 16):
804
  config_dim(c, args, "BLOCK_SIZE_M") == 1
805
  and getattr(args.get("Bs"), "dtype", None) == torch.float8_e4m3fn
806
  and args["M"] <= max_m
 
 
 
 
 
 
 
 
807
  )
808
 
809
  return config_filter(ok)
@@ -911,9 +987,9 @@ def swizzled_scale_config_pruner(allow_gate_subblock=False):
911
  def raise_no_swizzled_tile(configs, args):
912
  raise ValueError(
913
  "no autotune config can serve pre-swizzled scales for this launch (the "
914
- "SWIZZLE_32_4_4 read needs BLOCK_SIZE_K % 128 == 0 and a <=128-row or 128-multiple N tile; "
915
- f"GATE={bool(args.get('GATE'))}) — the contraction dim likely has no "
916
- "128-dividing tile; pass affine (row-major) scales for this shape."
917
  )
918
 
919
  return config_filter(
 
15
  import torch
16
  import triton
17
 
18
+ from .compat import FP8_DTYPE, get_active_device_type, is_sm10x, is_sm90, sm_count, sm_shared_memory_limit
19
  from .mma import MMA_N_ATOM_WIDTH
20
 
21
  # ── config pruners ────────────────────────────────────────────────────────────
 
101
  # it fit shared memory and failed only as benign launch-time smem overflows.
102
  SM10X_SCALED_MMA_MAX_N = 256
103
 
104
+ # Widest N tile the raw-activation (per-expert calibrated) arm survives; see
105
+ # raw_activation_pointer_pruner for the bisect.
106
+ RAW_ACT_MAX_BN = 128
107
+
108
  # Branch axes that RELOCATE the tile optimum (compute unit / operand orientation): the
109
  # tuner's guaranteed max-tile anchors group by these (``path_anchor_axes`` — a declaration,
110
  # so the tuner itself stays independent of configuration details). Scheduling axes
 
338
  dot/scalar/swap arms column-unpack them to E4M3 (lossless) first — no arm is
339
  structurally packed-incompatible, so W4A4 needs no shape gate of its own.
340
 
341
+ A weight-only launch whose raw activation is fp32 drops ``dot_scaled`` outright: the op's
342
+ lhs must be bf16/fp16, so the arm cannot compile at all there (arch-independent, unlike the
343
+ shape gates below).
344
+
345
  The shape gates above are scoped to ``dot_scaled`` — they are native scaled-MMA bug
346
  gates. The ``dot`` arm (BK structurally the UE8M0 group, 32) is CORRECT everywhere
347
  probed (forced-config sweep 2026-07-14, GATE and plain, MXFP4/MXFP8, incl. width-512
 
482
  return (2 if args.get("GATE") else 1) * config_dim(c, args, "BLOCK_SIZE_N") >= 128
483
  return config_dim(c, args, "BLOCK_SIZE_M") >= 128
484
 
485
+ def dot_scaled_act_dtype_ok(c, args):
486
+ # `tl.dot_scaled` takes a bf16/fp16 lhs. A weight-only launch hands it the RAW
487
+ # activation, so an fp32 one cannot reach the arm at all: every dot_scaled config dies
488
+ # in "Unexpected dtype for bf16. Got fp32" and the tuner books the whole arm as compile
489
+ # FAILURES rather than as a declared fence (87 of them on one grouped launch). A model
490
+ # running in fp32 reaches this through the transformers integration. The dot and scalar
491
+ # arms upcast the weight and take any float lhs, so they serve the launch — drop rather
492
+ # than raise. Not arch-gated: the lhs dtype is the op's contract, not an sm_10x quirk.
493
+ if c.kwargs.get("COMPUTE_MODE") != "dot_scaled" or acts_are_scaled(args):
494
+ return True
495
+ return getattr(args.get("A"), "dtype", None) != torch.float32
496
+
497
  def scales_are_e4m3(args):
498
  return getattr(args.get("Bs"), "dtype", None) == torch.float8_e4m3fn
499
 
 
515
  # re-admitting the single-trip/width traps it existed to remove).
516
  return compose_pruners(
517
  *stages,
518
+ config_filter(dot_scaled_act_dtype_ok),
519
  config_filter(nvfp4_native_ok, when=scales_are_e4m3),
520
  config_filter(
521
  mma_trap_ok, when=lambda args: is_sm10x(), on_empty=raise_all_mma_trapped
 
735
  return config_filter(ok)
736
 
737
 
738
+ # Pointer/pointer warp counts that trap on the WEIGHT-ONLY kernel under Triton < 3.8. Charted
739
+ # 2026-08-26 and REPRODUCED on 3.7.1 by the 2026-09-17 two-version sweep, so these rows are real
740
+ # measurements of that compiler, not the mis-chart the descriptor rows turned out to be. Triton
741
+ # 3.8 fixes every one of them, which is why `warp_spec_memory_mode_pruner` only consults the
742
+ # table below that version.
743
+ _WEIGHT_ONLY_POINTER_TRAPS = {
744
+ ("dot", "pointer"): {16},
745
+ ("dot_scaled", "pointer"): {4, 8, 16},
746
+ }
747
+
748
+
749
+ def _warp_spec_needs_descriptor_tile() -> bool:
750
+ """Triton >= 3.8, where a TMA-descriptor operand narrows WS to one tile (see the pruner)."""
751
+ return tuple(int(n) for n in triton.__version__.split(".")[:2]) >= (3, 8)
752
+
753
+
754
+ def warp_spec_memory_mode_pruner(weight_only: bool = False):
755
+ """``early_config_prune`` dropping ``warp_specialize`` configs the WS passes cannot lower
756
+ (``TritonGPUOptimizePartitionWarps`` / ``RelayoutTritonGPU`` -> ``PassManager::run failed``).
757
+
758
+ The rule is COMPILER-DEPENDENT, and the two versions are near mirror images. Charted
759
+ 2026-09-17 on B200 by filtering each kernel's OWN tuner configs (never hand-built — that skips
760
+ the descriptor pre_hooks and fails every descriptor cell for the wrong reason), BN held at 128,
761
+ BM and ``num_warps`` both axes, gpt-oss 2048x2048, on three kernels across two families:
762
+ grouped MX weight-only, grouped MX dynamic, and block-dynamic FP8.
763
+
764
+ operands triton 3.8.0 triton 3.7.1
765
+ pointer/pointer every warp count and BM mostly FAILS (see the table above)
766
+ desc/desc ONLY w4 AND BM == 128 most cells lower; w2 never does
767
+ mixed identical to desc/desc identical to desc/desc
768
+
769
+ Cell for cell across all three kernels, 3.8 rescues NO descriptor cell that 3.7.1 lowered
770
+ (0 of 96) and 3.7.1 rescues NO pointer cell that 3.8 lowered (0 of 80): descriptor+WS
771
+ REGRESSED in 3.8 and pointer+WS was FIXED. So each version gets the rule its own sweep
772
+ measured, rather than an intersection that would fence off a whole memory mode on both.
773
+
774
+ ``weight_only`` adds the rows that kernel charted for itself and the sweep did not re-probe:
775
+ ``BN == 32`` with WS, which failed 9/9 cells there, and (below 3.8) its trapping
776
+ pointer/pointer warp counts.
777
+
778
+ NOT covered here, deliberately: ``mx_dynamic``'s ``dot`` arm dies with a descriptor operand on
779
+ BOTH versions, and it is not a WS failure at all — it raises ``descriptor gather of uint8 must
780
+ have at least 32 columns, but got 16`` with ``warp_specialize`` on AND off. That is a gather
781
+ width law and belongs wherever gather width is decided."""
782
 
783
  def ok(c, args):
784
  if not c.kwargs.get("WARP_SPEC"):
785
  return True
786
+ if weight_only and config_dim(c, args, "BLOCK_SIZE_N") == 32:
 
 
787
  return False
788
  a, b = c.kwargs.get("A_MEMORY_MODE"), c.kwargs.get("B_MEMORY_MODE")
789
+ descriptor = a != "pointer" or b != "pointer"
790
+ if _warp_spec_needs_descriptor_tile():
791
+ # 3.8: one surviving descriptor cell, and pointer/pointer is wide open
792
+ return not descriptor or (c.num_warps == 4 and config_dim(c, args, "BLOCK_SIZE_M") == 128)
793
+ if descriptor:
794
  return True
795
+ traps = _WEIGHT_ONLY_POINTER_TRAPS if weight_only else {}
796
+ return c.num_warps not in traps.get((c.kwargs.get("COMPUTE_MODE"), a), set())
797
 
798
  return config_filter(ok, when=lambda args: is_sm10x())
799
 
 
813
  return config_filter(ok)
814
 
815
 
816
+ def raw_activation_pointer_pruner():
817
+ """``early_config_prune`` scoping the RAW-activation arm to the pointer load and away from
818
+ warp specialization. A calibrated scale held per expert leaves ``A`` unquantized for the
819
+ kernel to quantize per tile (a gathered row serves several experts, so there is no single
820
+ pre-quantized form), which puts a quantize INSIDE the K-loop and changes its structure:
821
+
822
+ - the TMA gather cannot serve it at all (``async_tma_gather`` wants 4 contiguous elements
823
+ per thread and the lowering fails), which at a deployment shape left the tuner no config;
824
+ - ``BLOCK_SIZE_N`` above 128 traps the device at prefill scale — a sticky misaligned
825
+ address. Bisected on B200 / Triton 3.8 with a COLD tune cache (a warm one replays a safe
826
+ crown and hides it): holding BN <= 128 runs clean, while fencing warp specialization, the
827
+ packed schedule, or BK < 128 each still trap. So the tile width is the variable, not the
828
+ schedule, the memory mode or WS.
829
+
830
+ A pre-quantized ``A`` keeps the full grid, unchanged — which is what the grouped prefill
831
+ now hands it (``quantize_routed_rows_per_expert`` lays the routed rows out quantized once
832
+ the row count pays for the copy), so the raw arm is reached below that regime, or above it
833
+ when ``FINEGRAINED_FORCE_GATHER`` holds the gather. The BN fence is load-bearing there."""
834
+
835
+ def ok(c, args):
836
+ a = args.get("A")
837
+ if getattr(a, "dtype", None) == FP8_DTYPE or args.get("As") is None:
838
+ return True
839
+ return (
840
+ c.kwargs.get("A_MEMORY_MODE", "pointer") == "pointer"
841
+ and config_dim(c, args, "BLOCK_SIZE_N") <= RAW_ACT_MAX_BN
842
+ )
843
+
844
+ return config_filter(ok)
845
+
846
+
847
  def mx_2d_swap_scope_pruner(max_m: int = 16):
848
  """``early_config_prune`` scoping the 2D mx kernel's ``SWAP_AB`` rows to the ONE regime
849
  they exist for: E4M3-scale (NVFP4) single-token decode. NVFP4 has no other native M=1
 
872
  config_dim(c, args, "BLOCK_SIZE_M") == 1
873
  and getattr(args.get("Bs"), "dtype", None) == torch.float8_e4m3fn
874
  and args["M"] <= max_m
875
+ # An A-side TMA descriptor under SWAP_AB traps the device on Triton 3.8 — a sticky
876
+ # misaligned address, after which every later launch reports it wherever it came
877
+ # from. Observed on B200 in a GLM-5.2-NVFP4 forward (16 tokens; clean at 4) at
878
+ # dot_scaled, BM=1, BN ∈ {128, 256}, BK ∈ {64, 128}: the tile WIDTH is not the
879
+ # variable (BN=128 traps too), the A descriptor is. Narrowed by probe — fencing the
880
+ # B side as well was NOT needed, so the weight descriptor stays available to the
881
+ # crown, and pointer-mode A keeps the decode arm this pruner exists for.
882
+ and c.kwargs.get("A_MEMORY_MODE", "pointer") == "pointer"
883
  )
884
 
885
  return config_filter(ok)
 
987
  def raise_no_swizzled_tile(configs, args):
988
  raise ValueError(
989
  "no autotune config can serve pre-swizzled scales for this launch (the "
990
+ "SWIZZLE_32_4_4 read needs BLOCK_SIZE_K % 128 == 0 and a <=128-row or 128-multiple N "
991
+ f"tile; N={args.get('N')}, K={args.get('K')}, GATE={bool(args.get('GATE'))}) — say "
992
+ "which dim is short rather than guessing: pass affine (row-major) scales for it."
993
  )
994
 
995
  return config_filter(
build/torch-rocm/quant.py CHANGED
@@ -19,7 +19,7 @@ from triton.language.extra.cuda import gdc_launch_dependents, gdc_wait
19
 
20
  from .bayesian_autotuner import bayesian_autotune
21
  from .formats import global_scale_stride, is_per_expert_global
22
- from .compat import add_op_namespace_prefix, FP8_DTYPE, MX_SCALE_GROUP_K, NVFP4_SCALE_GROUP_K, compile_time_only_triton_op, compile_time_only_triton_wrap, decode_pdl, device_context, is_sm10x
23
  from .swizzle import swizzle_store_block
24
  from .scheduling import build_tile_layout, resolve_tile_inline
25
 
@@ -927,9 +927,121 @@ def _fp8_act_quant_kernel(
927
  tl.store(s_ptr + pid, s)
928
 
929
 
930
- @compile_time_only_triton_op(
931
- add_op_namespace_prefix("fp8_act_quant_tensor_wide"), mutates_args=(), opaque=True
932
- )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
933
  def fp8_act_quant_tensor_wide(
934
  x: torch.Tensor, block_size: int = 128
935
  ) -> tuple[torch.Tensor, torch.Tensor]:
 
19
 
20
  from .bayesian_autotuner import bayesian_autotune
21
  from .formats import global_scale_stride, is_per_expert_global
22
+ from .compat import FP8_DTYPE, MX_SCALE_GROUP_K, NVFP4_SCALE_GROUP_K, compile_time_only_triton_wrap, decode_pdl, device_context, is_sm10x
23
  from .swizzle import swizzle_store_block
24
  from .scheduling import build_tile_layout, resolve_tile_inline
25
 
 
927
  tl.store(s_ptr + pid, s)
928
 
929
 
930
+ def static_expert_act_operands(A, As, num_experts):
931
+ """Activation operands of the per-expert STATIC arms: ``(A_q, As, as_stride)``.
932
+
933
+ One calibrated scale for the whole matmul, or one per expert — a MoE calibrates each
934
+ separately. ``as_stride`` 0 makes every expert read the one scale, 1 gives each its own.
935
+ Per expert the kernel quantizes in register against the tile's own scale, because one token
936
+ routes to several experts whose scales differ and so has no single pre-quantized form; one
937
+ scale for all of them pre-quantizes once here instead.
938
+ """
939
+ As = As.reshape(-1).to(torch.float32)
940
+ as_stride = int(As.numel() == num_experts)
941
+ A_q = A if as_stride else (A.to(torch.float32) / As).to(FP8_DTYPE)
942
+ return A_q, As, as_stride
943
+
944
+
945
+ # Host-side operand marshalling, NOT a custom op: it returns its inputs unchanged on the
946
+ # pre-quantized and per-expert branches, which `torch.library.custom_op` forbids (the output may
947
+ # not alias an input), and it hands back Python strides. The quant kernels it calls carry their
948
+ # own compile handling.
949
+ def tensor_wide_act_operands(
950
+ A: torch.Tensor, As: torch.Tensor | None, num_experts: int | None = None
951
+ ) -> tuple[torch.Tensor, torch.Tensor, int, int]:
952
+ """Activation operands of the tensor-scale FP8 arms: ``(A, As, stride_m, stride_e)``.
953
+
954
+ No ``As`` derives one scale per row here; a pre-quantized (E4M3) ``A`` brings its own. An
955
+ ``As`` on a raw ``A`` is the CALIBRATED (static) scale that replaces them: one value, which
956
+ ``A`` quantizes against here so every row reads the one entry, or — where the op routes —
957
+ one per expert, which stays raw for the kernel to quantize each tile against its own expert
958
+ — a gathered row serves several experts, so it has no single pre-quantized form until the
959
+ routed rows are laid out (``quantize_routed_rows_per_expert``). ``num_experts`` ``None`` on
960
+ the 2D op, whose rows belong to no expert.
961
+ """
962
+ if As is None:
963
+ A, As = fp8_act_quant_tensor_wide(A, A.shape[-1])
964
+ As = As.reshape(-1)
965
+ return A, As, As.stride(0), 0
966
+ if A.dtype == FP8_DTYPE:
967
+ return A, As, As.stride(0), 0
968
+ As = As.reshape(-1).float()
969
+ assert As.numel() in (1, num_experts), (
970
+ f"a calibrated activation scale is one value, or one per expert where the op routes; "
971
+ f"got {As.numel()}"
972
+ )
973
+ if As.numel() == 1:
974
+ return quantize_rows_static(A, As), As, 0, 0
975
+ return A, As, 0, 1
976
+
977
+
978
+ @triton.jit
979
+ def _static_row_scale(S, ExpertStart, row, stride_s, NUM_EXPERTS: tl.constexpr, PER_EXPERT: tl.constexpr):
980
+ """This row's calibrated scale. ``PER_EXPERT`` resolves it from the routing itself — the
981
+ rows are expert-sorted, so the count of expert ends at or below the row IS its expert,
982
+ which costs one ``(NUM_EXPERTS,)`` compare and spares the host materializing a per-row
983
+ vector (a gather or a repeat_interleave, either of which outweighs this whole kernel)."""
984
+ if PER_EXPERT:
985
+ ends = tl.load(ExpertStart + 1 + tl.arange(0, NUM_EXPERTS))
986
+ s = tl.load(S + tl.sum((ends <= row).to(tl.int32)))
987
+ else:
988
+ s = tl.load(S + row * stride_s)
989
+ return s
990
+
991
+
992
+ @triton.jit
993
+ def _fp8_quant_static_kernel(
994
+ X, Y, S, GatherIdx, ExpertStart, stride_x_m, stride_y_m, stride_s, K,
995
+ NUM_EXPERTS: tl.constexpr, BLOCK_SIZE_K: tl.constexpr,
996
+ GATHER: tl.constexpr, PER_EXPERT: tl.constexpr,
997
+ ):
998
+ """One row tile against a PROVIDED scale: gather, divide in fp32, store E4M3. The fp32
999
+ divide is the in-register static arm's arithmetic, so both forms of ``A`` round alike.
1000
+ ``stride_s`` 0 broadcasts one scale over every row; ``GATHER`` reads ``X`` through
1001
+ ``GatherIdx`` so the routed rows land laid out."""
1002
+ row = tl.program_id(0)
1003
+ src = tl.load(GatherIdx + row) if GATHER else row
1004
+ offs = tl.program_id(1) * BLOCK_SIZE_K + tl.arange(0, BLOCK_SIZE_K)
1005
+ mask = offs < K
1006
+ x = tl.load(X + src.to(tl.int64) * stride_x_m + offs, mask=mask, other=0.0)
1007
+ y = (x.to(tl.float32) / _static_row_scale(S, ExpertStart, row, stride_s, NUM_EXPERTS, PER_EXPERT)).to(tl.float8e4nv)
1008
+ tl.store(Y + row * stride_y_m + offs, y, mask=mask)
1009
+
1010
+
1011
+ def quantize_rows_static(A, scales, gather_idx=None, expert_start=None, num_experts=None):
1012
+ """``A`` as E4M3 against ``scales`` — one per row, one for all of them, or (with
1013
+ ``expert_start``) one per expert resolved from the routing. ``gather_idx`` lays the routed
1014
+ rows out in the same pass. One read and one write, where the torch spelling costs four."""
1015
+ assert A.stride(-1) == 1, "the static row quant reads K contiguous"
1016
+ rows = gather_idx.shape[0] if gather_idx is not None else A.shape[0]
1017
+ K = A.shape[-1]
1018
+ y = A.new_empty(rows, K, dtype=FP8_DTYPE)
1019
+ block_k = min(triton.next_power_of_2(K), 1024)
1020
+ with device_context(A.device):
1021
+ compile_time_only_triton_wrap(_fp8_quant_static_kernel)[(rows, triton.cdiv(K, block_k))](
1022
+ A, y, scales, gather_idx, expert_start,
1023
+ A.stride(0), y.stride(0), 0 if scales.numel() == 1 else scales.stride(0), K,
1024
+ NUM_EXPERTS=num_experts or 1, BLOCK_SIZE_K=block_k,
1025
+ GATHER=gather_idx is not None, PER_EXPERT=expert_start is not None,
1026
+ )
1027
+ return y
1028
+
1029
+
1030
+ def quantize_routed_rows_per_expert(A, As, gather_idx, expert_start, num_experts):
1031
+ """The routed rows as E4M3, each quantized against its own expert's calibrated scale.
1032
+
1033
+ A per-expert scale otherwise hands the kernel raw rows to quantize per tile, which reads
1034
+ ``A`` at bf16 width through every N-tile and holds the tile at ``RAW_ACT_MAX_BN``. Laying
1035
+ the rows out once buys back both, at the rounding the in-register arm produces, because a
1036
+ row's expert is known once the rows are expert-sorted — ``gather_idx`` orders them, or the
1037
+ intermediate already is. ``As`` is untouched: the kernel dequantizes per expert whichever
1038
+ form ``A`` arrives in.
1039
+
1040
+ Returns the laid-out rows; the caller's ``gather_idx`` is spent by that layout.
1041
+ """
1042
+ return quantize_rows_static(A, As, gather_idx, expert_start, num_experts)
1043
+
1044
+
1045
  def fp8_act_quant_tensor_wide(
1046
  x: torch.Tensor, block_size: int = 128
1047
  ) -> tuple[torch.Tensor, torch.Tensor]: