543 lines
27 KiB
Diff
543 lines
27 KiB
Diff
diff --git a/tools/validate_cute_nvfp4_h3.py b/tools/validate_cute_nvfp4_h3.py
|
|
new file mode 100644
|
|
index 0000000..f7a3aa3
|
|
--- /dev/null
|
|
+++ b/tools/validate_cute_nvfp4_h3.py
|
|
@@ -0,0 +1,537 @@
|
|
+"""Compare the CUTLASS DSL SM121 block-scaled GEMM with real H3 NVFP4 tensors."""
|
|
+
|
|
+from __future__ import annotations
|
|
+
|
|
+import argparse
|
|
+import importlib.util
|
|
+import json
|
|
+import sys
|
|
+from pathlib import Path
|
|
+
|
|
+import torch
|
|
+import torch.nn.functional as functional
|
|
+
|
|
+from h3_blackwell_runtime.nvfp4_quant import vortex_quantize_nvfp4
|
|
+from profile_nvfp4_linear import module_for_name, representative_inputs
|
|
+
|
|
+
|
|
+def _replace_once(source: str, old: str, new: str, name: str) -> str:
|
|
+ if source.count(old) != 1:
|
|
+ raise RuntimeError(
|
|
+ f"CUTLASS {name} patch expected one occurrence, found {source.count(old)}: {old!r}"
|
|
+ )
|
|
+ return source.replace(old, new)
|
|
+
|
|
+
|
|
+def _patch_streaming_a(source: str) -> str:
|
|
+ source = _replace_once(
|
|
+ source,
|
|
+ "import cutlass\nimport cutlass.cute as cute",
|
|
+ """import cutlass
|
|
+import cutlass.cute as cute
|
|
+from cutlass import Float32
|
|
+from cutlass.cutlass_dsl import dsl_user_op
|
|
+from cutlass._mlir import ir
|
|
+from cutlass._mlir.dialects import llvm
|
|
+
|
|
+
|
|
+@dsl_user_op
|
|
+def vortex_rcp_approx_ftz_f32(
|
|
+ x: Float32,
|
|
+ *,
|
|
+ loc: Optional[ir.Location] = None,
|
|
+ ip: Optional[ir.InsertionPoint] = None,
|
|
+) -> Float32:
|
|
+ result = llvm.inline_asm(
|
|
+ Float32.mlir_type,
|
|
+ [x.ir_value(loc=loc, ip=ip)],
|
|
+ "rcp.approx.ftz.f32 $0, $1;",
|
|
+ "=f,f",
|
|
+ has_side_effects=False,
|
|
+ asm_dialect=0,
|
|
+ loc=loc,
|
|
+ ip=ip,
|
|
+ )
|
|
+ return Float32(result)""",
|
|
+ "streaming-A reciprocal",
|
|
+ )
|
|
+ replacements = (
|
|
+ (
|
|
+ " self.a_dtype = a.element_type\n self.b_dtype = b.element_type\n self.c_dtype = c.element_type\n self.sf_dtype = sfa.element_type",
|
|
+ " self.a_source_dtype = a.element_type\n self.a_dtype = cutlass.Float4E2M1FN\n self.b_dtype = b.element_type\n self.c_dtype = c.element_type\n self.sf_dtype = cutlass.Float8E4M3FN",
|
|
+ ),
|
|
+ (
|
|
+ " self.sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(\n a.shape, self.sf_vec_size\n )\n sfa_tensor = cute.make_tensor(sfa.iterator, self.sfa_layout)\n",
|
|
+ "",
|
|
+ ),
|
|
+ (
|
|
+ " tma_atom_a, tma_tensor_a = self._make_tma_atoms_and_tensors(\n a,\n self.a_smem_layout_staged,\n (self.tile_shape_mnk[0], self.tile_shape_mnk[2]),\n 1,\n internal_type=self.tma_internal_a_dtype,\n )\n\n",
|
|
+ "",
|
|
+ ),
|
|
+ (
|
|
+ " tma_atom_sfa, tma_tensor_sfa = self._make_tma_atoms_and_tensors(\n sfa_tensor,\n self.sfa_smem_layout_staged,\n (self.tile_shape_mnk[0], self.tile_shape_mnk[2]),\n 1,\n internal_type=cutlass.Int16,\n )\n\n",
|
|
+ "",
|
|
+ ),
|
|
+ (
|
|
+ " self.kernel(\n tma_atom_a,\n tma_tensor_a,\n tma_atom_b,\n tma_tensor_b,\n tma_atom_sfa,\n tma_tensor_sfa,\n tma_atom_sfb,",
|
|
+ " self.kernel(\n a,\n sfa,\n tma_atom_b,\n tma_tensor_b,\n tma_atom_sfb,",
|
|
+ ),
|
|
+ (
|
|
+ " tma_atom_a: cute.CopyAtom,\n mA_mkl: cute.Tensor,\n tma_atom_b: cute.CopyAtom,\n mB_nkl: cute.Tensor,\n tma_atom_sfa: cute.CopyAtom,\n mSFA_mkl: cute.Tensor,\n tma_atom_sfb: cute.CopyAtom,",
|
|
+ " mA_mkl: cute.Tensor,\n tensor_scale_a: cute.Tensor,\n tma_atom_b: cute.CopyAtom,\n mB_nkl: cute.Tensor,\n tma_atom_sfb: cute.CopyAtom,",
|
|
+ ),
|
|
+ (
|
|
+ " cpasync.prefetch_descriptor(tma_atom_a)\n cpasync.prefetch_descriptor(tma_atom_b)\n cpasync.prefetch_descriptor(tma_atom_sfa)\n cpasync.prefetch_descriptor(tma_atom_sfb)",
|
|
+ " cpasync.prefetch_descriptor(tma_atom_b)\n cpasync.prefetch_descriptor(tma_atom_sfb)",
|
|
+ ),
|
|
+ (
|
|
+ " tma_copy_bytes = (\n cute.size_in_bytes(self.a_dtype, a_smem_layout)\n + cute.size_in_bytes(self.b_dtype, b_smem_layout)\n + cute.size_in_bytes(self.sf_dtype, sfa_smem_layout)\n + cute.size_in_bytes(self.sf_dtype, sfb_smem_layout)\n )",
|
|
+ " tma_copy_bytes = (\n cute.size_in_bytes(self.b_dtype, b_smem_layout)\n + cute.size_in_bytes(self.sf_dtype, sfb_smem_layout)\n )",
|
|
+ ),
|
|
+ (
|
|
+ " # (tM, tK, loopM, loopK, loopL)\n gSFA_mkl = cute.local_tile(\n mSFA_mkl,\n cute.slice_(self.tile_shape_mnk, (None, 0, None)),\n (None, None, None),\n )\n",
|
|
+ "",
|
|
+ ),
|
|
+ (
|
|
+ " # TMA load A partition_S/D\n a_cta_layout = cute.make_layout(cute.slice_(cta_layout_mnk, (0, None, 0)).shape)\n a_cta_crd = cluster_coord_mnk[1]\n tAsA, tAgA = cpasync.tma_partition(\n tma_atom_a,\n a_cta_crd,\n a_cta_layout,\n cute.group_modes(sA, 0, 2),\n cute.group_modes(gA_mkl, 0, 2),\n )\n\n",
|
|
+ "",
|
|
+ ),
|
|
+ (
|
|
+ " tAsSFA, tAgSFA = cpasync.tma_partition(\n tma_atom_sfa,\n a_cta_crd,\n a_cta_layout,\n cute.group_modes(sSFA, 0, 2),\n cute.group_modes(gSFA_mkl, 0, 2),\n )\n tAsSFA = cute.filter_zeros(tAsSFA)\n tAgSFA = cute.filter_zeros(tAgSFA)\n\n",
|
|
+ "",
|
|
+ ),
|
|
+ (
|
|
+ " cute.arch.setmaxregister_decrease(self.load_register_requirement)\n\n while work_tile.is_valid_tile:",
|
|
+ """ cute.arch.setmaxregister_decrease(self.load_register_requirement)
|
|
+ producer_lane = tidx - self.tma_load_warp_id * self.num_threads_per_warp
|
|
+ fp4_store = cute.make_copy_atom(
|
|
+ cute.nvgpu.CopyUniversalOp(), cutlass.Float4E2M1FN
|
|
+ )
|
|
+ source_fragment = cute.make_rmem_tensor((16,), cutlass.Float32)
|
|
+ normalized_fragment = cute.make_rmem_tensor((8,), cutlass.Float32)
|
|
+ fp4_fragment = cute.make_rmem_tensor((8,), cutlass.Float4E2M1FN)
|
|
+ scale_source = cute.make_rmem_tensor((8,), cutlass.Float32)
|
|
+ scale_fragment = cute.make_rmem_tensor((8,), cutlass.Float8E4M3FN)
|
|
+ decoded_scale_fragment = cute.make_rmem_tensor((8,), cutlass.Float32)
|
|
+ while work_tile.is_valid_tile:""",
|
|
+ ),
|
|
+ (
|
|
+ " tAgA_mkl = tAgA[(None, tile_coord_mnl[0], None, tile_coord_mnl[2])]\n tBgB_nkl = tBgB[(None, tile_coord_mnl[1], None, tile_coord_mnl[2])]\n tAgSFA_mkl = tAgSFA[(None, tile_coord_mnl[0], None, tile_coord_mnl[2])]\n tBgSFB_nkl = tBgSFB[(None, tile_coord_mnl[1], None, tile_coord_mnl[2])]",
|
|
+ " tBgB_nkl = tBgB[(None, tile_coord_mnl[1], None, tile_coord_mnl[2])]\n tBgSFB_nkl = tBgSFB[(None, tile_coord_mnl[1], None, tile_coord_mnl[2])]",
|
|
+ ),
|
|
+ (
|
|
+ " tAgA_k = tAgA_mkl[(None, mainloop_producer_state.count)]\n tAsA_pipe = tAsA[(None, mainloop_producer_state.index)]\n\n tBgB_k = tBgB_nkl[(None, mainloop_producer_state.count)]\n tBsB_pipe = tBsB[(None, mainloop_producer_state.index)]\n\n tAgSFA_k = tAgSFA_mkl[(None, mainloop_producer_state.count)]\n tAsSFA_pipe = tAsSFA[(None, mainloop_producer_state.index)]\n\n tBgSFB_k = tBgSFB_nkl[(None, mainloop_producer_state.count)]",
|
|
+ " tBgB_k = tBgB_nkl[(None, mainloop_producer_state.count)]\n tBsB_pipe = tBsB[(None, mainloop_producer_state.index)]\n\n tBgSFB_k = tBgSFB_nkl[(None, mainloop_producer_state.count)]",
|
|
+ ),
|
|
+ (
|
|
+ """ cute.copy(
|
|
+ tma_atom_a,
|
|
+ tAgA_k,
|
|
+ tAsA_pipe,
|
|
+ tma_bar_ptr=mainloop_pipeline.producer_get_barrier(
|
|
+ mainloop_producer_state
|
|
+ ),
|
|
+ )
|
|
+""",
|
|
+ """ scale = tensor_scale_a[0]
|
|
+ stage = mainloop_producer_state.index
|
|
+ source_k_base = mainloop_producer_state.count * self.tile_shape_mnk[2]
|
|
+ source_m_base = tile_coord_mnl[0] * self.tile_shape_mnk[0]
|
|
+ for row_group in cutlass.range_constexpr(4):
|
|
+ row = producer_lane + row_group * self.num_threads_per_warp
|
|
+ source_row = source_m_base + row
|
|
+ sA_row = sA[(row, None, stage)]
|
|
+ sA_tiles = cute.zipped_divide(sA_row, (8,))
|
|
+ for block_column in cutlass.range_constexpr(8):
|
|
+ source_column = source_k_base + block_column * 16
|
|
+ maximum = cutlass.Float32(0.0)
|
|
+ for element in cutlass.range_constexpr(16):
|
|
+ value = mA_mkl[
|
|
+ source_row,
|
|
+ source_column + element,
|
|
+ tile_coord_mnl[2],
|
|
+ ]
|
|
+ source_fragment[element] = value
|
|
+ maximum = cutlass.max(
|
|
+ cutlass.max(value, -value), maximum
|
|
+ )
|
|
+
|
|
+ raw_block_scale = (
|
|
+ maximum / cutlass.Float32(6.0)
|
|
+ ) / scale
|
|
+ for element in cutlass.range_constexpr(8):
|
|
+ scale_source[element] = raw_block_scale
|
|
+ scale_values = scale_source.load()
|
|
+ scale_values = cute.where(
|
|
+ scale_values <= cutlass.Float32(448.0),
|
|
+ scale_values,
|
|
+ cutlass.Float32(448.0),
|
|
+ )
|
|
+ scale_fragment.store(
|
|
+ scale_values.to(cutlass.Float8E4M3FN)
|
|
+ )
|
|
+ sSFA[row, block_column * 16, stage] = scale_fragment[0]
|
|
+ decoded_scale_fragment.store(
|
|
+ scale_fragment.load().to(cutlass.Float32)
|
|
+ )
|
|
+ raw_encode_scale = vortex_rcp_approx_ftz_f32(
|
|
+ decoded_scale_fragment[0] * scale
|
|
+ )
|
|
+ for element in cutlass.range_constexpr(8):
|
|
+ scale_source[element] = raw_encode_scale
|
|
+ encode_scale_values = scale_source.load()
|
|
+ encode_scale_values = cute.where(
|
|
+ encode_scale_values
|
|
+ <= cutlass.Float32(3.402823466e38),
|
|
+ encode_scale_values,
|
|
+ cutlass.Float32(3.402823466e38),
|
|
+ )
|
|
+ scale_source.store(encode_scale_values)
|
|
+ encode_scale = scale_source[0]
|
|
+
|
|
+ for half in cutlass.range_constexpr(2):
|
|
+ for element in cutlass.range_constexpr(8):
|
|
+ normalized_fragment[element] = source_fragment[
|
|
+ half * 8 + element
|
|
+ ] * encode_scale
|
|
+ fp4_fragment.store(
|
|
+ normalized_fragment.load().to(
|
|
+ cutlass.Float4E2M1FN
|
|
+ )
|
|
+ )
|
|
+ cute.copy(
|
|
+ fp4_store,
|
|
+ fp4_fragment,
|
|
+ sA_tiles[(None, block_column * 2 + half)],
|
|
+ )
|
|
+ cute.arch.fence_proxy("async.shared", space="cta")
|
|
+""",
|
|
+ ),
|
|
+ (
|
|
+ """ cute.copy(
|
|
+ tma_atom_sfa,
|
|
+ tAgSFA_k,
|
|
+ tAsSFA_pipe,
|
|
+ tma_bar_ptr=mainloop_pipeline.producer_get_barrier(
|
|
+ mainloop_producer_state
|
|
+ ),
|
|
+ )
|
|
+""",
|
|
+ "",
|
|
+ ),
|
|
+ )
|
|
+ for index, (old, new) in enumerate(replacements):
|
|
+ source = _replace_once(source, old, new, f"streaming-A[{index}]")
|
|
+ return source
|
|
+
|
|
+
|
|
+def load_cutlass_example(path: Path, *, fuse_alpha: bool, stream_a: bool = False):
|
|
+ if not path.is_file():
|
|
+ raise FileNotFoundError(f"CUTLASS DSL example not found: {path}")
|
|
+ sys.path.insert(0, str(path.parent))
|
|
+ source = path.read_text(encoding="utf-8")
|
|
+ if stream_a:
|
|
+ source = _patch_streaming_a(source)
|
|
+ if fuse_alpha:
|
|
+ replacements = (
|
|
+ (" c: cute.Tensor,\n max_active_clusters: cutlass.Constexpr,", " c: cute.Tensor,\n alpha: cute.Tensor,\n max_active_clusters: cutlass.Constexpr,"),
|
|
+ (" mC_mnl: cute.Tensor,\n tiled_mma: cute.TiledMma,", " mC_mnl: cute.Tensor,\n alpha: cute.Tensor,\n tiled_mma: cute.TiledMma,"),
|
|
+ (" tma_tensor_c,\n self.tiled_mma,", " tma_tensor_c,\n alpha,\n self.tiled_mma,"),
|
|
+ (" tRS_rD_out.store(acc_vec.to(self.c_dtype))", " tRS_rD_out.store((acc_vec * alpha[0]).to(self.c_dtype))"),
|
|
+ )
|
|
+ for old, new in replacements:
|
|
+ source = _replace_once(source, old, new, "alpha")
|
|
+ if stream_a or fuse_alpha:
|
|
+ suffix = "_vortex_stream_a" if stream_a else "_vortex"
|
|
+ suffix += "_alpha" if fuse_alpha else ""
|
|
+ load_path = path.with_name(f"{path.stem}{suffix}.py")
|
|
+ load_path.write_text(source, encoding="utf-8")
|
|
+ else:
|
|
+ load_path = path
|
|
+ spec = importlib.util.spec_from_file_location("vortex_cutlass_blockscaled", load_path)
|
|
+ if spec is None or spec.loader is None:
|
|
+ raise ImportError(f"Cannot load CUTLASS DSL example: {load_path}")
|
|
+ module = importlib.util.module_from_spec(spec)
|
|
+ spec.loader.exec_module(module)
|
|
+ return module
|
|
+
|
|
+
|
|
+def decode_comfy_fp4(storage: torch.Tensor) -> torch.Tensor:
|
|
+ lookup = torch.tensor(
|
|
+ [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0, -0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0],
|
|
+ device=storage.device,
|
|
+ dtype=torch.float32,
|
|
+ )
|
|
+ codes = torch.stack((storage >> 4, storage & 0x0F), dim=-1).reshape(storage.shape[0], -1)
|
|
+ return lookup[codes.long()]
|
|
+
|
|
+
|
|
+def fp4_tensor(storage: torch.Tensor, *, swap_nibbles: bool, reencode: bool):
|
|
+ import cutlass
|
|
+ import cutlass.torch as cutlass_torch
|
|
+ from cutlass.cute.runtime import from_dlpack
|
|
+
|
|
+ rows, packed_columns = storage.shape
|
|
+ if reencode:
|
|
+ logical = decode_comfy_fp4(storage).unsqueeze(-1)
|
|
+ return cutlass_torch.cute_tensor_like(
|
|
+ logical, cutlass.Float4E2M1FN, is_dynamic_layout=True, assumed_align=16,
|
|
+ )
|
|
+ if swap_nibbles:
|
|
+ storage = ((storage & 0x0F) << 4) | ((storage & 0xF0) >> 4)
|
|
+ # DLPack cannot export Torch's packed FP4 dtype yet. Build the desired
|
|
+ # logical layout from an oversized uint8 allocation, then reinterpret its
|
|
+ # iterator as FP4 and populate only the packed storage that the layout uses.
|
|
+ backing = torch.empty(
|
|
+ (rows, packed_columns * 2, 1), device=storage.device, dtype=torch.uint8,
|
|
+ )
|
|
+ backing.zero_()
|
|
+ backing[:, :packed_columns, 0].copy_(storage)
|
|
+ tensor = from_dlpack(backing, assumed_align=16)
|
|
+ tensor.element_type = cutlass.Float4E2M1FN
|
|
+ tensor = tensor.mark_layout_dynamic(leading_dim=1)
|
|
+ return tensor, backing
|
|
+
|
|
+
|
|
+def output_tensor(storage: torch.Tensor):
|
|
+ from cutlass.cute.runtime import from_dlpack
|
|
+
|
|
+ tensor = from_dlpack(storage.unsqueeze(-1), assumed_align=16)
|
|
+ tensor = tensor.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=1)
|
|
+ return tensor
|
|
+
|
|
+
|
|
+def scale_tensor(storage: torch.Tensor):
|
|
+ import cutlass
|
|
+ from cutlass.cute.runtime import from_dlpack
|
|
+
|
|
+ tensor = from_dlpack(storage.view(torch.uint8).unsqueeze(-1), assumed_align=16)
|
|
+ tensor.element_type = cutlass.Float8E4M3FN
|
|
+ return tensor.mark_layout_dynamic(leading_dim=1)
|
|
+
|
|
+
|
|
+def parse_args() -> argparse.Namespace:
|
|
+ parser = argparse.ArgumentParser(description=__doc__)
|
|
+ parser.add_argument("--cutlass-example", type=Path, required=True)
|
|
+ parser.add_argument("--model-path", default="/models/minimax_h3_fl2va_pruned_nvfp4.safetensors")
|
|
+ parser.add_argument("--output", type=Path, required=True)
|
|
+ parser.add_argument("--linear", choices=("attn_qkv_proj", "attn_out_proj", "mlp_fc1", "mlp_fc2"), default="attn_qkv_proj")
|
|
+ parser.add_argument("--rows", type=int, default=128)
|
|
+ parser.add_argument("--tile-k", type=int, choices=(128, 256), default=128)
|
|
+ parser.add_argument("--block-index", type=int, default=24)
|
|
+ parser.add_argument("--width", type=int, default=1344)
|
|
+ parser.add_argument("--height", type=int, default=768)
|
|
+ parser.add_argument("--frames", type=int, default=124)
|
|
+ parser.add_argument("--steps", type=int, default=12)
|
|
+ parser.add_argument("--sampler-step", type=int, default=1)
|
|
+ parser.add_argument("--seed", type=int, default=440420)
|
|
+ parser.add_argument("--text-tokens", type=int, default=100)
|
|
+ parser.add_argument("--attention", default="sage2")
|
|
+ parser.add_argument("--device", default="cuda")
|
|
+ parser.add_argument("--swap-nibbles", action="store_true")
|
|
+ parser.add_argument("--reencode-fp4", action="store_true")
|
|
+ parser.add_argument("--zero-a", action="store_true")
|
|
+ parser.add_argument("--fuse-alpha", action="store_true")
|
|
+ parser.add_argument("--stream-a", action="store_true")
|
|
+ parser.add_argument("--warmup", type=int, default=0)
|
|
+ parser.add_argument("--iterations", type=int, default=0)
|
|
+ parser.add_argument("--benchmark-reference", action="store_true")
|
|
+ return parser.parse_args()
|
|
+
|
|
+
|
|
+def main() -> None:
|
|
+ args = parse_args()
|
|
+ if args.rows <= 0 or args.rows % 128:
|
|
+ raise ValueError("--rows must be a positive multiple of 128")
|
|
+ if args.stream_a and args.linear == "mlp_fc2":
|
|
+ raise ValueError("--stream-a deliberately excludes mlp_fc2; retain the reference cuBLAS path")
|
|
+ if args.warmup < 0 or args.iterations < 0:
|
|
+ raise ValueError("--warmup and --iterations must be non-negative")
|
|
+
|
|
+ import cutlass
|
|
+ import cutlass.cute as cute
|
|
+ import cutlass.torch as cutlass_torch
|
|
+ import comfy_kitchen as ck
|
|
+ from cutlass.cute.runtime import from_dlpack
|
|
+ from comfy_kitchen.tensor import TensorCoreNVFP4Layout
|
|
+
|
|
+ example = load_cutlass_example(
|
|
+ args.cutlass_example, fuse_alpha=args.fuse_alpha, stream_a=args.stream_a,
|
|
+ )
|
|
+ block, inputs, metadata = representative_inputs(args)
|
|
+ linear = module_for_name(block, args.linear)
|
|
+ activation = inputs[args.linear].reshape(-1, linear.in_features)[:args.rows].contiguous()
|
|
+
|
|
+ with torch.inference_mode():
|
|
+ packed_activation = vortex_quantize_nvfp4(activation)
|
|
+ packed_weight = linear._packed_weight()
|
|
+ a_qdata, a_tensor_scale, a_block_scale = TensorCoreNVFP4Layout.get_plain_tensors(packed_activation)
|
|
+ b_qdata, b_tensor_scale, b_block_scale = TensorCoreNVFP4Layout.get_plain_tensors(packed_weight)
|
|
+ reference = functional.linear(packed_activation, packed_weight, None)[:args.rows, :linear.out_features]
|
|
+ raw_reference = ck.scaled_mm_nvfp4(
|
|
+ a_qdata,
|
|
+ b_qdata,
|
|
+ tensor_scale_a=a_tensor_scale,
|
|
+ tensor_scale_b=b_tensor_scale,
|
|
+ block_scale_a=a_block_scale,
|
|
+ block_scale_b=b_block_scale,
|
|
+ out_dtype=torch.bfloat16,
|
|
+ alpha=torch.ones(1, device=activation.device, dtype=torch.float32),
|
|
+ )[:args.rows, :linear.out_features]
|
|
+ if args.zero_a:
|
|
+ a_qdata = torch.zeros_like(a_qdata)
|
|
+
|
|
+ output_bf16 = torch.zeros((a_qdata.shape[0], b_qdata.shape[0]), device=activation.device, dtype=torch.bfloat16)
|
|
+ if args.stream_a:
|
|
+ if args.rows != 128 or activation.shape[1] % 128:
|
|
+ raise ValueError("--stream-a currently requires exactly 128 rows and K divisible by 128")
|
|
+ a = output_tensor(activation)
|
|
+ a_backing = activation
|
|
+ else:
|
|
+ a, a_backing = fp4_tensor(a_qdata, swap_nibbles=args.swap_nibbles, reencode=args.reencode_fp4)
|
|
+ b, b_backing = fp4_tensor(b_qdata, swap_nibbles=args.swap_nibbles, reencode=args.reencode_fp4)
|
|
+ sfa = (
|
|
+ from_dlpack(a_tensor_scale.float().reshape(1).contiguous(), assumed_align=4)
|
|
+ if args.stream_a
|
|
+ else scale_tensor(a_block_scale)
|
|
+ )
|
|
+ sfb = scale_tensor(b_block_scale)
|
|
+ c = output_tensor(output_bf16)
|
|
+
|
|
+ gemm = example.Sm120BlockScaledGemmKernel(
|
|
+ cutlass.Float32,
|
|
+ 16,
|
|
+ (128, 128, args.tile_k),
|
|
+ (128, 128),
|
|
+ )
|
|
+ hardware_info = cutlass.utils.HardwareInfo()
|
|
+ max_active_clusters = hardware_info.get_max_active_clusters(1)
|
|
+ stream = cutlass_torch.default_stream()
|
|
+ alpha = a_tensor_scale.float() * b_tensor_scale.float()
|
|
+ alpha_argument = from_dlpack(alpha.reshape(1).contiguous(), assumed_align=4)
|
|
+ if args.fuse_alpha:
|
|
+ compiled = cute.compile(gemm, a, b, sfa, sfb, c, alpha_argument, max_active_clusters, stream)
|
|
+ arguments = (a, b, sfa, sfb, c, alpha_argument, stream)
|
|
+ else:
|
|
+ compiled = cute.compile(gemm, a, b, sfa, sfb, c, max_active_clusters, stream)
|
|
+ arguments = (a, b, sfa, sfb, c, stream)
|
|
+ compiled(*arguments)
|
|
+ torch.cuda.synchronize()
|
|
+
|
|
+ timing = None
|
|
+ if args.iterations:
|
|
+ for _ in range(args.warmup):
|
|
+ compiled(*arguments)
|
|
+ torch.cuda.synchronize()
|
|
+ started = torch.cuda.Event(enable_timing=True)
|
|
+ finished = torch.cuda.Event(enable_timing=True)
|
|
+ started.record()
|
|
+ for _ in range(args.iterations):
|
|
+ compiled(*arguments)
|
|
+ finished.record()
|
|
+ finished.synchronize()
|
|
+ total_ms = started.elapsed_time(finished)
|
|
+ timing = {
|
|
+ "warmup": args.warmup,
|
|
+ "iterations": args.iterations,
|
|
+ "total_ms": total_ms,
|
|
+ "mean_ms": total_ms / args.iterations,
|
|
+ }
|
|
+
|
|
+ reference_timing = None
|
|
+ if args.benchmark_reference:
|
|
+ if not args.iterations:
|
|
+ raise ValueError("--benchmark-reference requires --iterations")
|
|
+
|
|
+ def measure_cuda(fn):
|
|
+ result = None
|
|
+ for _ in range(args.warmup):
|
|
+ result = fn()
|
|
+ torch.cuda.synchronize()
|
|
+ started = torch.cuda.Event(enable_timing=True)
|
|
+ finished = torch.cuda.Event(enable_timing=True)
|
|
+ started.record()
|
|
+ for _ in range(args.iterations):
|
|
+ result = fn()
|
|
+ finished.record()
|
|
+ finished.synchronize()
|
|
+ total = started.elapsed_time(finished)
|
|
+ return result, total / args.iterations
|
|
+
|
|
+ _, quantize_ms = measure_cuda(lambda: vortex_quantize_nvfp4(activation))
|
|
+ _, complete_ms = measure_cuda(
|
|
+ lambda: functional.linear(
|
|
+ vortex_quantize_nvfp4(activation), packed_weight, None,
|
|
+ )
|
|
+ )
|
|
+ reference_timing = {
|
|
+ "backend": "vortex_scale_plus_comfy_pack_gemm",
|
|
+ "activation_quantize_mean_ms": quantize_ms,
|
|
+ "complete_projection_mean_ms": complete_ms,
|
|
+ }
|
|
+
|
|
+ candidate = (
|
|
+ output_bf16[:args.rows, :linear.out_features]
|
|
+ if args.fuse_alpha
|
|
+ else (output_bf16[:args.rows, :linear.out_features].float() * alpha).to(reference.dtype)
|
|
+ )
|
|
+ raw_candidate = output_bf16[:args.rows, :linear.out_features]
|
|
+ raw_delta = raw_candidate.float() - raw_reference.float()
|
|
+ delta = candidate.float() - reference.float()
|
|
+ report = {
|
|
+ "device": torch.cuda.get_device_name(),
|
|
+ "torch": torch.__version__,
|
|
+ "cutlass_dsl": "4.6.2",
|
|
+ "metadata": metadata,
|
|
+ "linear": args.linear,
|
|
+ "rows": args.rows,
|
|
+ "mnk": [args.rows, linear.out_features, linear.in_features],
|
|
+ "tile_shape_mnk": [128, 128, args.tile_k],
|
|
+ "swap_nibbles": args.swap_nibbles,
|
|
+ "reencode_fp4": args.reencode_fp4,
|
|
+ "zero_a": args.zero_a,
|
|
+ "fuse_alpha": args.fuse_alpha,
|
|
+ "stream_a": args.stream_a,
|
|
+ "streamed_activation_materialization": (
|
|
+ {"global_qdata": False, "global_sfa": False}
|
|
+ if args.stream_a
|
|
+ else None
|
|
+ ),
|
|
+ "timing": timing,
|
|
+ "reference_timing": reference_timing,
|
|
+ "cute_shapes": {"a": str(a.shape), "b": str(b.shape), "sfa": str(sfa.shape), "sfb": str(sfb.shape), "c": str(c.shape)},
|
|
+ "tensor_scales": {"a": a_tensor_scale.float().item(), "b": b_tensor_scale.float().item(), "alpha": alpha.item()},
|
|
+ "raw_output": {
|
|
+ "dtype": str(output_bf16.dtype),
|
|
+ "checksum": output_bf16.float().sum().item(),
|
|
+ "max_abs": output_bf16.float().abs().max().item(),
|
|
+ "nonzero": int(torch.count_nonzero(output_bf16).item()),
|
|
+ "finite": bool(torch.isfinite(output_bf16).all().item()),
|
|
+ },
|
|
+ "raw_blockscaled_parity": {
|
|
+ "applicable": not args.fuse_alpha,
|
|
+ "reference_checksum": raw_reference.float().sum().item(),
|
|
+ "candidate_checksum": raw_candidate.float().sum().item(),
|
|
+ "equal": torch.equal(raw_candidate, raw_reference),
|
|
+ "max_abs": raw_delta.abs().max().item(),
|
|
+ "mean_abs": raw_delta.abs().mean().item(),
|
|
+ },
|
|
+ "reference_checksum": reference.float().sum().item(),
|
|
+ "candidate_checksum": candidate.float().sum().item(),
|
|
+ "equal": torch.equal(candidate, reference),
|
|
+ "max_abs": delta.abs().max().item(),
|
|
+ "mean_abs": delta.abs().mean().item(),
|
|
+ "relative_l2": (delta.norm() / reference.float().norm().clamp_min(1e-12)).item(),
|
|
+ "numerical_note": (
|
|
+ "The experimental epilogue applies the FP32 global-scale product before BF16 conversion."
|
|
+ if args.fuse_alpha
|
|
+ else "The stock SM121 kernel rounds before the external global-scale product; exact H3 integration requires the fused-alpha epilogue."
|
|
+ ),
|
|
+ }
|
|
+ args.output.parent.mkdir(parents=True, exist_ok=True)
|
|
+ args.output.write_text(json.dumps(report, indent=2) + "\n", encoding="utf-8")
|
|
+ print(json.dumps(report, indent=2), flush=True)
|
|
+
|
|
+
|
|
+if __name__ == "__main__":
|
|
+ main()
|