Kernels:
Trusted publisher
Uploaded using `kernel-builder`.
Browse files- build/torch-rocm/_ops.py +1 -1
- build/torch-rocm/batched.py +79 -66
- build/torch-rocm/bayesian_autotuner.py +6 -5
- build/torch-rocm/compat.py +8 -0
- build/torch-rocm/grouped.py +49 -27
- build/torch-rocm/loading/tiles.py +12 -9
- build/torch-rocm/matmul.py +71 -20
- build/torch-rocm/metadata.json +12 -12
- build/torch-rocm/metadata.json.sigstore +1 -1
- build/torch-rocm/moe.py +53 -14
- build/torch-rocm/pruners.py +116 -40
- build/torch-rocm/quant.py +116 -4
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 = "
|
| 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,
|
| 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
|
| 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, #
|
| 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 |
-
|
|
|
|
|
|
|
| 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, #
|
| 416 |
-
B, # (num_experts, N, K) FP8
|
| 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 |
-
|
|
|
|
|
|
|
| 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,
|
| 484 |
for _ in range(0, tl.cdiv(K, BLOCK_SIZE_K)):
|
| 485 |
-
a, _ =
|
| 486 |
b, _ = load_weight_plain(
|
| 487 |
-
b_ptrs, b_ptrs, 0, 0, 0,
|
| 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 |
-
|
| 497 |
-
#
|
| 498 |
-
|
| 499 |
-
|
| 500 |
-
|
| 501 |
if PDL:
|
| 502 |
gdc_launch_dependents()
|
| 503 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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,)
|
|
|
|
| 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 =
|
| 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
|
|
|
|
| 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,
|
|
|
|
|
|
|
| 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 |
-
|
| 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 |
-
|
|
|
|
| 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 |
-
|
| 1823 |
-
|
| 1824 |
-
|
| 1825 |
-
|
| 1826 |
-
|
|
|
|
|
|
|
| 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,
|
|
|
|
| 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 |
-
|
| 105 |
-
|
| 106 |
-
|
| 107 |
-
|
| 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,
|
| 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,
|
| 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, #
|
| 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,
|
| 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, #
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 =
|
| 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 |
-
|
| 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,)
|
|
|
|
| 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 =
|
| 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
|
|
|
|
| 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 |
-
|
| 1599 |
-
|
| 1600 |
-
#
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
| 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 |
-
|
| 2238 |
-
|
| 2239 |
-
|
| 2240 |
-
|
| 2241 |
-
|
|
|
|
|
|
|
| 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
|
| 272 |
-
|
| 273 |
-
|
| 274 |
-
|
| 275 |
-
|
| 276 |
-
|
| 277 |
-
|
| 278 |
-
|
| 279 |
-
|
|
|
|
|
|
|
|
|
|
| 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,
|
| 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 |
-
#
|
| 1283 |
-
qA, As =
|
| 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 |
-
|
| 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.
|
| 1378 |
-
|
| 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
|
| 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("
|
| 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": "
|
| 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": "
|
| 15 |
"backward.py": "ri9lKPhb9SmkYBvUS8KvuORiW5PcKViNamFNVWxb69g=",
|
| 16 |
-
"batched.py": "
|
| 17 |
-
"bayesian_autotuner.py": "
|
| 18 |
-
"compat.py": "
|
| 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": "
|
| 24 |
"loading/__init__.py": "so7kYmlVoiASfoOzemdU5fBTiP8yA85gVPNQExIq/KU=",
|
| 25 |
"loading/scales.py": "8HJpflmImEXYJquA6hVLStUUEecvExAxBEYVb3qPPWE=",
|
| 26 |
-
"loading/tiles.py": "
|
| 27 |
-
"matmul.py": "
|
| 28 |
"mma.py": "1fg2an7Dc8Vd6qpSmq+Q0781A6sNilOJHU0PqEbouT8=",
|
| 29 |
-
"moe.py": "
|
| 30 |
"norm.py": "VvE6GUE0Oztf7PDcviLCcB4fhnXt7/zgdE+aApv6DBM=",
|
| 31 |
-
"pruners.py": "
|
| 32 |
-
"quant.py": "
|
| 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": "
|
| 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
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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)
|
|
|
|
|
|
|
|
|
|
| 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``.
|
| 566 |
-
|
|
|
|
|
|
|
| 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 |
-
|
| 669 |
-
|
|
|
|
| 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
|
| 823 |
-
each GEMM quantizes its raw input against its
|
|
|
|
|
|
|
| 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 |
-
|
| 718 |
-
|
| 719 |
-
|
| 720 |
-
|
| 721 |
-
|
| 722 |
-
|
| 723 |
-
|
| 724 |
-
|
| 725 |
-
|
| 726 |
-
|
| 727 |
-
|
| 728 |
-
|
| 729 |
-
|
| 730 |
-
|
| 731 |
-
|
| 732 |
-
|
| 733 |
-
|
| 734 |
-
|
| 735 |
-
|
| 736 |
-
|
| 737 |
-
|
| 738 |
-
|
| 739 |
-
|
| 740 |
-
|
| 741 |
-
|
| 742 |
-
|
| 743 |
-
|
| 744 |
-
|
| 745 |
-
|
| 746 |
-
|
| 747 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 748 |
|
| 749 |
def ok(c, args):
|
| 750 |
if not c.kwargs.get("WARP_SPEC"):
|
| 751 |
return True
|
| 752 |
-
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 758 |
return True
|
| 759 |
-
|
|
|
|
| 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
|
| 915 |
-
f"GATE={bool(args.get('GATE'))}) —
|
| 916 |
-
"
|
| 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
|
| 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 |
-
|
| 931 |
-
|
| 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]:
|