Release v4.0.0 (#2294)

This commit is contained in:
Kihiro Bando
2025-05-13 15:55:29 -04:00
committed by GitHub
parent ad7b2f5e84
commit f115c3f854
299 changed files with 51495 additions and 4413 deletions
@@ -0,0 +1,392 @@
# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: BSD-3-Clause
# Redistribution and use in source and binary forms, with or without
# modification, are permitted provided that the following conditions are met:
# 1. Redistributions of source code must retain the above copyright notice, this
# list of conditions and the following disclaimer.
# 2. Redistributions in binary form must reproduce the above copyright notice,
# this list of conditions and the following disclaimer in the documentation
# and/or other materials provided with the distribution.
# 3. Neither the name of the copyright holder nor the names of its
# contributors may be used to endorse or promote products derived from
# this software without specific prior written permission.
# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
import argparse
import torch
import time
from typing import Type
import cuda.bindings.driver as cuda
import cutlass
import cutlass.cute as cute
from cutlass.cute.runtime import from_dlpack
import cutlass.torch as cutlass_torch
"""
An Elementwise Addition Example using CuTe DSL.
This example kernel copies data from global memory to register memory (rmem), performs the elementwise
addition operation, and stores the result back to global memory.
Primary goals of this example are to demonstrate how basic global memory copies can be expressed in
CuTe DSL and illustrate canonical partitioning patterns in CuTe. It also implements canonical
predication for tensors whose shape is not multiple of tile size to guard OOB reads.
Thread-value (or TV) layouts are central to canonical partitioning patterns in CuTe. They provide a
mapping from thread and a thread's value to the set of coordinates within a tile that we have sliced
out from a data tensor.
The input tensors are row-major layout, that leading dimension is the right most dimension. In order
to efficiently copy data from global memory, we must map threads contiguously on row dimension.
Thread ID mapping to 2D coordinates with layout `(4,32):(32,1)`:
+----+----+----+----+-----+----+
| | 0 | 1 | 2 | ... | 31 |
+----+----+----+----+-----+----+
| 0 | T0 | T1 | T2 | ... | T31|
+----+----+----+----+-----+----+
| 1 |T32 |T33 |T34 | ... |T63 |
+----+----+----+----+-----+----+
| 2 |T64 |T65 |T66 | ... |T95 |
+----+----+----+----+-----+----+
| 3 |T96 |T97 |T98 | ... |T127|
+----+----+----+----+-----+----+
As Ampere GPU supports a maximum of 128bit per load/store instruction and each element is 32bit, we
can load 4 elements per instruction. Having additional contiguous values allows for vectorization
across threads (coalesced accesses) and is required for saturating the memory bandwidth.
We use `(4,4):(4,1)` as the val layout in this example. Notice that the major mode is the same as
the major mode of the input tensor - without which vectorization would not be possible.
If you already know the TV layout you want to use for your tiled copy, CuTe DSL provides utility
`cute.make_layout_tv` to build the tiled copy type around it and the atom of your choice.
.. code-block:: python
thr_layout = cute.make_layout((4, 32), stride=(32, 1))
val_layout = cute.make_layout((4, 4), stride=(4, 1))
tiler_mn, tv_layout = cute.make_layout_tv(thr_layout, val_layout)
# Tile input tensor to thread blocks: ((TileM,TileN),(RestM,RestN))
gA = cute.zipped_divide(mA, tiler_mn)
where `tiler_mn` is the tile size per thread block and `tv_layout` is the TV layout which maps
thread index and inter-thread index of data array per thread to logical coordinates of elements in
input and output tensors.
Then we can build tiled copy for input and output tensors with `cute.make_tiled_copy` utility.
.. code-block:: python
blkA = gA[((None, None), bidx)] # (TileM,TileN)
copy_atom_load = cute.make_copy_atom(cute.nvgpu.CopyUniversalOp(), gA.element_type)
tiled_copy_A = cute.make_tiled_copy(copy_atom_load, tv_layout, tiler_mn)
# get slice of tiled_copy_A for current thread
thr_copy_A = tiled_copy_A.get_slice(tidx)
# partition per thread block tensor as source of tiled copy
thrA = thr_copy_A.partition_S(blkA)
# allocate fragment for gmem->rmem
frgA = cute.make_fragment_like(thrA)
# copy data from global memory to register memory
cute.copy(copy_atom_load, thrA, frgA)
To run this example:
.. code-block:: bash
python examples/ampere/elementwise_add.py --M 3 --N 12
python examples/ampere/elementwise_add.py --M 1024 --N 512
python examples/ampere/elementwise_add.py --M 1024 --N 1024 --benchmark --warmup_iterations 2 --iterations 1000
To collect performance with NCU profiler:
.. code-block:: bash
# Don't iterate too many times when profiling with ncu
ncu python examples/ampere/elementwise_add.py --M 2048 --N 2048 --benchmark --iterations 10 --skip_ref_check
"""
@cute.kernel
def elementwise_add_kernel(
gA: cute.Tensor,
gB: cute.Tensor,
gC: cute.Tensor,
cC: cute.Tensor, # coordinate tensor
shape: cute.Shape,
tv_layout: cute.Layout,
tiler_mn: cute.Shape,
):
tidx, _, _ = cute.arch.thread_idx()
bidx, _, _ = cute.arch.block_idx()
# slice for CTAs
# logical id -> address
blk_coord = ((None, None), bidx)
blkA = gA[blk_coord] # (TileM,TileN)
blkB = gB[blk_coord] # (TileM,TileN)
blkC = gC[blk_coord] # (TileM,TileN)
blkCrd = cC[blk_coord] # (TileM, TileN)
print(f"[DSL INFO] Sliced Tensors per thread block:")
print(f"[DSL INFO] blkA = {blkA.type}")
print(f"[DSL INFO] blkB = {blkB.type}")
print(f"[DSL INFO] blkC = {blkC.type}")
print(f"[DSL INFO] blkCrd = {blkCrd.type}")
# # declare the atoms which will be used later for memory copy
copy_atom_load = cute.make_copy_atom(cute.nvgpu.CopyUniversalOp(), gA.element_type)
copy_atom_store = cute.make_copy_atom(cute.nvgpu.CopyUniversalOp(), gC.element_type)
tiled_copy_A = cute.make_tiled_copy(copy_atom_load, tv_layout, tiler_mn)
tiled_copy_B = cute.make_tiled_copy(copy_atom_load, tv_layout, tiler_mn)
tiled_copy_C = cute.make_tiled_copy(copy_atom_store, tv_layout, tiler_mn)
thr_copy_A = tiled_copy_A.get_slice(tidx)
thr_copy_B = tiled_copy_B.get_slice(tidx)
thr_copy_C = tiled_copy_C.get_slice(tidx)
thrA = thr_copy_A.partition_S(blkA)
thrB = thr_copy_B.partition_S(blkB)
thrC = thr_copy_C.partition_S(blkC)
# allocate fragments for gmem->rmem
frgA = cute.make_fragment_like(thrA)
frgB = cute.make_fragment_like(thrB)
frgC = cute.make_fragment_like(thrC)
thrCrd = thr_copy_C.partition_S(blkCrd)
frgPred = cute.make_fragment(thrCrd.shape, cutlass.Boolean)
print(f"[DSL INFO] Sliced Tensors per thread:")
print(f"[DSL INFO] thrA = {thrA.type}")
print(f"[DSL INFO] thrB = {thrB.type}")
print(f"[DSL INFO] thrC = {thrC.type}")
print(f"[DSL INFO] thrCrd = {thrCrd.type}")
for i in cutlass.range_dynamic(0, cute.size(frgPred), 1):
val = cute.elem_less(thrCrd[i], shape)
frgPred[i] = val
# Print per thread predicate mask
# if tidx == 0 and bidx == 0:
# cute.printf("block_dim = {}", cute.arch.grid_dim())
# cute.printf("shape = {}", shape)
# cute.print_tensor(thrA)
# cute.print_tensor(thrB)
# cute.print_tensor(frgPred)
##########################################################
# Move data to reg address space
##########################################################
cute.copy(copy_atom_load, thrA, frgA, pred=frgPred)
cute.copy(copy_atom_load, thrB, frgB, pred=frgPred)
# if tidx == 0 and bidx == 0:
# cute.print_tensor(frgA)
# cute.print_tensor(frgB)
# Load data before use. The compiler will optimize the copy and load
# operations to convert some memory ld/st into register uses.
result = frgA.load() + frgB.load()
# Save the results back to registers. Here we reuse b's registers.
frgC.store(result)
# Copy the results back to c
cute.copy(copy_atom_store, frgC, thrC, pred=frgPred)
@cute.jit
def elementwise_add(mA, mB, mC, copy_bits: cutlass.Constexpr = 128):
dtype = mA.element_type
vector_size = copy_bits // dtype.width
thr_layout = cute.make_ordered_layout((4, 32), order=(1, 0))
val_layout = cute.make_ordered_layout((4, vector_size), order=(1, 0))
tiler_mn, tv_layout = cute.make_layout_tv(thr_layout, val_layout)
print(f"[DSL INFO] Input Tensors:")
print(f"[DSL INFO] mA = {mA.type}")
print(f"[DSL INFO] mB = {mB.type}")
print(f"[DSL INFO] Tiling Parameters:")
print(f"[DSL INFO] tiler_mn = {tiler_mn} per thread block")
print(f"[DSL INFO] tv_layout = {tv_layout}")
gA = cute.zipped_divide(mA, tiler_mn) # ((TileM,TileN),(RestM,RestN))
gB = cute.zipped_divide(mB, tiler_mn) # ((TileM,TileN),(RestM,RestN))
gC = cute.zipped_divide(mC, tiler_mn) # ((TileM,TileN),(RestM,RestN))
print(f"[DSL INFO] Tiled Tensors:")
print(f"[DSL INFO] gA = {gA.type}")
print(f"[DSL INFO] gB = {gB.type}")
print(f"[DSL INFO] gC = {gC.type}")
idC = cute.make_identity_tensor(mC.shape)
cC = cute.zipped_divide(idC, tiler=tiler_mn)
print(f"[DSL INFO] coord tensor = {cC.type}")
elementwise_add_kernel(gA, gB, gC, cC, mC.shape, tv_layout, tiler_mn).launch(
grid=[cute.size(gC, mode=[1]), 1, 1],
block=[cute.size(tv_layout, mode=[0]), 1, 1],
)
def run_elementwise_add(
M,
N,
dtype: Type[cutlass.Numeric],
is_a_dynamic_layout=False,
is_b_dynamic_layout=False,
is_result_dynamic_layout=False,
skip_ref_check=False,
benchmark=True,
warmup_iterations=2,
iterations=200,
):
if not torch.cuda.is_available():
raise RuntimeError(f"Ampere GPU is required to run this example!")
print(f"\nRunning Elementwise Add test with:")
print(f"Tensor dimensions: [{M}, {N}]")
print(f"Input and Output Data type: {dtype}")
torch_dtype = cutlass_torch.dtype(dtype)
if dtype.is_integer:
a = torch.randint(0, 10, (M, N), device=torch.device("cuda"), dtype=torch_dtype)
b = torch.randint(0, 10, (M, N), device=torch.device("cuda"), dtype=torch_dtype)
else:
a = torch.randn(M, N, device=torch.device("cuda"), dtype=torch_dtype)
b = torch.randn(M, N, device=torch.device("cuda"), dtype=torch_dtype)
c = torch.zeros_like(a)
print(f"Input tensor shapes:")
print(f"a: {a.shape}, dtype: {a.dtype}")
print(f"b: {b.shape}, dtype: {b.dtype}")
print(f"c: {c.shape}, dtype: {c.dtype}\n")
if not is_a_dynamic_layout:
a_tensor = from_dlpack(a).mark_layout_dynamic()
else:
a_tensor = a
if not is_b_dynamic_layout:
b_tensor = from_dlpack(b).mark_layout_dynamic()
else:
b_tensor = b
if not is_result_dynamic_layout:
c_tensor = from_dlpack(c).mark_layout_dynamic()
else:
c_tensor = c
print("Compiling kernel with cute.compile ...")
start_time = time.time()
compiled_func = cute.compile(elementwise_add, a_tensor, b_tensor, c_tensor)
compilation_time = time.time() - start_time
print(f"Compilation time: {compilation_time:.4f} seconds")
print("Executing vector add kernel...")
# Get current CUDA stream from PyTorch
torch_stream = torch.cuda.current_stream()
# Get the raw stream pointer as a CUstream
current_stream = cuda.CUstream(torch_stream.cuda_stream)
if not skip_ref_check:
compiled_func(a_tensor, b_tensor, c_tensor)
print("Verifying results...")
torch.testing.assert_close(a + b, c)
print("Results verified successfully!")
if not benchmark:
return
# Create CUDA events for timing
start_event = cuda.cuEventCreate(cuda.CUevent_flags.CU_EVENT_DEFAULT)[1]
end_event = cuda.cuEventCreate(cuda.CUevent_flags.CU_EVENT_DEFAULT)[1]
# Warmup
for _ in range(warmup_iterations):
compiled_func(a_tensor, b_tensor, c_tensor)
# Use the current stream for CUDA events instead of the default stream
# Record start event
cuda.cuEventRecord(start_event, current_stream)
# Execute the kernel
for _ in range(iterations):
compiled_func(a_tensor, b_tensor, c_tensor)
# Record end event
cuda.cuEventRecord(end_event, current_stream)
cuda.cuEventSynchronize(end_event)
# Calculate elapsed time
err, elapsed_time = cuda.cuEventElapsedTime(start_event, end_event)
avg_time = elapsed_time / iterations
# Print execution results
print(f"Kernel execution time: {avg_time:.4f} ms")
print(
f"Achieved memory throughput: {(3 * a.numel() * dtype.width // 8) / (avg_time / 1000) / 1e9:.2f} GB/s"
)
print(f"First few elements of result: \n{c[:3, :3]}")
# Destroy events
cuda.cuEventDestroy(start_event)
cuda.cuEventDestroy(end_event)
if __name__ == "__main__":
parser = argparse.ArgumentParser(
description="example of elementwise add to demonstrate the numpy/pytorch as input for kernels"
)
parser.add_argument("--M", default=1024, type=int)
parser.add_argument("--N", default=1024, type=int)
parser.add_argument("--warmup_iterations", default=2, type=int)
parser.add_argument("--iterations", default=100, type=int)
parser.add_argument("--skip_ref_check", action="store_true")
parser.add_argument("--benchmark", action="store_true")
args = parser.parse_args()
run_elementwise_add(
args.M,
args.N,
dtype=cutlass.Float32,
is_a_dynamic_layout=True,
is_b_dynamic_layout=True,
is_result_dynamic_layout=True,
skip_ref_check=args.skip_ref_check,
benchmark=args.benchmark,
warmup_iterations=args.warmup_iterations,
iterations=args.iterations,
)
print("\nPASS")
@@ -0,0 +1,395 @@
# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: BSD-3-Clause
# Redistribution and use in source and binary forms, with or without
# modification, are permitted provided that the following conditions are met:
# 1. Redistributions of source code must retain the above copyright notice, this
# list of conditions and the following disclaimer.
# 2. Redistributions in binary form must reproduce the above copyright notice,
# this list of conditions and the following disclaimer in the documentation
# and/or other materials provided with the distribution.
# 3. Neither the name of the copyright holder nor the names of its
# contributors may be used to endorse or promote products derived from
# this software without specific prior written permission.
# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
import argparse
import operator
import torch
from typing import Type
import time
import cuda.bindings.driver as cuda
import cutlass
import cutlass.cute as cute
import cutlass.torch as cutlass_torch
from cutlass.cute.runtime import from_dlpack
"""
An Elementwise Apply Example using CuTe DSL.
This example kernel demonstrates the meta-programming capability of the CuTe DSL by allowing
customization of elementwise operations through lambda functions. The kernel copies data from
global memory to register memory (rmem), applies a user-defined operation to the elements,
and stores the result back to global memory.
Primary goals of this example:
1. Demonstrate meta-programming capability by passing lambda functions to customize elementwise operations
2. Show how to apply different operations (add, multiply, etc.) using the same kernel structure
3. Illustrate how to parameterize CUDA kernels with operation types at compile time
To run this example:
.. code-block:: bash
# Run with addition operation
python examples/ampere/elementwise_apply.py --M 1024 --N 512 --op add
# Run with multiplication operation
python examples/ampere/elementwise_apply.py --M 1024 --N 512 --op mul
# Run with subtraction operation
python examples/ampere/elementwise_apply.py --M 1024 --N 512 --op sub
# Benchmark performance
python examples/ampere/elementwise_apply.py --M 2048 --N 2048 --op add --benchmark --warmup_iterations 2 --iterations 10
The example demonstrates how to express complex CUDA kernels with customizable operations
while maintaining high performance through efficient memory access patterns.
"""
@cute.kernel
def elementwise_apply_kernel(
op: cutlass.Constexpr,
gA: cute.Tensor,
gB: cute.Tensor,
gC: cute.Tensor,
cC: cute.Tensor, # coordinate tensor
shape: cute.Shape,
tv_layout: cute.Layout, # (tid, vid) -> logic coord
):
tidx, _, _ = cute.arch.thread_idx()
bidx, _, _ = cute.arch.block_idx()
# slice for CTAs
cta_coord = ((None, None), bidx)
# logical coord -> address
ctaA = gA[cta_coord] # (TileM, TileN)
ctaB = gB[cta_coord] # (TileM, TileN)
ctaC = gC[cta_coord] # (TileM, TileN)
ctaCrd = cC[cta_coord] # (TileM, TileN)
print(f"[DSL INFO] Sliced Tensors per thread block:")
print(f"[DSL INFO] ctaA = {ctaA.type}")
print(f"[DSL INFO] ctaB = {ctaB.type}")
print(f"[DSL INFO] ctaC = {ctaC.type}")
print(f"[DSL INFO] ctaCrd = {ctaCrd.type}")
# compose with CTA TV layout
# (tid, vid) -> address
tidfrgA = cute.composition(ctaA, tv_layout)
tidfrgB = cute.composition(ctaB, tv_layout)
tidfrgC = cute.composition(ctaC, tv_layout)
tidfrgCrd = cute.composition(ctaCrd, tv_layout)
# print(f"{tv_layout = }")
# print(f"{tidfrgA = }")
thr_coord = (tidx, (None, None))
# slice for threads
# vid -> address
thrA = tidfrgA[thr_coord] # (V)
thrB = tidfrgB[thr_coord] # (V)
thrC = tidfrgC[thr_coord] # (V)
thrCrd = tidfrgCrd[thr_coord]
print(f"[DSL INFO] Sliced Tensors per thread:")
print(f"[DSL INFO] thrA = {thrA.type}")
print(f"[DSL INFO] thrB = {thrB.type}")
print(f"[DSL INFO] thrC = {thrC.type}")
print(f"[DSL INFO] thrCrd = {thrCrd.type}")
# allocate fragments for gmem->rmem
frgA = cute.make_fragment_like(thrA, gA.element_type)
frgB = cute.make_fragment_like(thrB, gB.element_type)
frgC = cute.make_fragment_like(thrC, gC.element_type)
frgPred = cute.make_fragment(thrCrd.shape, cutlass.Boolean)
for i in cutlass.range_dynamic(cute.size(frgPred), unroll=1):
frgPred[i] = cute.elem_less(thrCrd[i], shape)
# if tidx == 0 and bidx == 0:
# cute.print_tensor(frgPred)
##########################################################
# Move data to reg address space
##########################################################
# declare the atoms which will be used later for memory copy
copy_atom_load = cute.make_copy_atom(
cute.nvgpu.CopyUniversalOp(),
gA.element_type,
num_bits_per_copy=gA.element_type.width,
)
copy_atom_store = cute.make_copy_atom(
cute.nvgpu.CopyUniversalOp(),
gC.element_type,
num_bits_per_copy=gC.element_type.width,
)
cute.copy(copy_atom_load, thrA, frgA, pred=frgPred)
cute.copy(copy_atom_load, thrB, frgB, pred=frgPred)
# Load data before use. The compiler will optimize the copy and load
# operations to convert some memory ld/st into register uses.
result = op(frgA.load(), frgB.load())
# Save the results back to registers. Here we reuse b's registers.
frgC.store(result)
# Copy the results back to c
cute.copy(copy_atom_store, frgC, thrC, pred=frgPred)
@cute.jit
def elementwise_apply(
op: cutlass.Constexpr,
a: cute.Tensor,
b: cute.Tensor,
result: cute.Tensor,
):
"""CUDA kernel applying binary operator on each element of two n-D input tensors in
CuTe Python and store to result tensor.
:param op: Binary operator or lambda function to apply element-wise
:type op: cutlass.Constexpr
:param a: First input tensor
:type a: cute.Tensor
:param b: Second input tensor
:type b: cute.Tensor
:param result: Output tensor to store the results of op(a, b)
:type result: cute.Tensor
:return: None
:rtype: None
.. code-block:: python
# Example 1: Adding two tensors
x = torch.tensor([[1, 2], [3, 4]], dtype=torch.float32, device="cuda")
y = torch.tensor([[5, 6], [7, 8]], dtype=torch.float32, device="cuda")
result = torch.empty_like(x)
elementwise_apply(operator.add, from_dlpack(x), from_dlpack(y), from_dlpack(result))
# result:
# tensor([[6.0, 8.0],
# [10.0, 12.0]], device='cuda:0')
# Example 2: Using a lambda function
elementwise_apply(lambda a, b: a * a + b * b, from_dlpack(x), from_dlpack(y), from_dlpack(result))
# result:
# tensor([[ 2., 8.],
# [ 54., 512.]], device='cuda:0')
"""
# Baseline: naive TV layout
# * mA layout: (4096, 4096):(4096, 1)
# * TV layout map to (512, 4) tile
# * tidx maps to mode-0 but input layout is contiguous on mode-1, performance will be bad
# tv_layout = cute.make_layout((128, (4, 4)), stride=(4, (512, 1)))
# cta_tiler = (512, 4)
# Opt-1: better TV layout with better 1D thread layout (SOL with 1D thread layout)
# * mA layout: (4096, 4096):(4096, 1)
# * TV layout map to (4, 512) tile
# * tidx maps to mode-1 which is leading mode of input tensor for coalesced load
# tv_layout = cute.make_layout((128, (4, 4)), stride=(16, (4, 1)))
# cta_tiler = (4, 512)
# Opt-2: 2D tile but worse
# * mA layout: (4096, 4096):(4096, 1)
# * TV layout map to (128, 16) logical tile
# * V layout is bad as contiguous mode is not on right-most
# * `cute.copy` only supports vectorize when stride-1 of v-layout on right-most )
# tv_layout = cute.make_layout(((32, 4), (4, 4)), stride=((4, 512), (1, 128)))
# cta_tiler = (128, 16)
# Opt-3: SOL with 2D thread tile
# * mA layout: (4096, 4096):(4096, 1)
# * TV layout map to (16, 128) logical tile
# * tidx maps to mode-1 and input layout is contiguous on mode-1 for coalesced load-store
thr_layout = cute.make_layout((4, 32), stride=(32, 1))
val_layout = cute.make_layout((4, 4), stride=(4, 1))
tiler_mn, tv_layout = cute.make_layout_tv(thr_layout, val_layout)
print(f"[DSL INFO] Input Tensors:")
print(f"[DSL INFO] a = {a.type}")
print(f"[DSL INFO] b = {b.type}")
print(f"[DSL INFO] result = {result.type}")
print(f"[DSL INFO] Tiling Parameters:")
print(f"[DSL INFO] tiler_mn = {tiler_mn} per thread block")
print(f"[DSL INFO] tv_layout = {tv_layout}")
gA = cute.zipped_divide(a, tiler_mn) # ((TileM, TileN), (RestM, RestN))
gB = cute.zipped_divide(b, tiler_mn) # ((TileM, TileN), (RestM, RestN))
gC = cute.zipped_divide(result, tiler_mn) # ((TileM, TileN), (RestM, RestN))
print(f"[DSL INFO] Tiled Tensors:")
print(f"[DSL INFO] gA = {gA.type}")
print(f"[DSL INFO] gB = {gB.type}")
print(f"[DSL INFO] gC = {gC.type}")
idC = cute.make_identity_tensor(result.shape)
cC = cute.zipped_divide(idC, tiler=tiler_mn)
print(f"[DSL INFO] coord tensor = {cC.type}")
# Launch the kernel asynchronously
# Async token(s) can also be specified as dependencies
elementwise_apply_kernel(
op,
gA,
gB,
gC,
cC,
result.shape,
tv_layout,
).launch(
grid=[cute.size(gC, mode=[1]), 1, 1],
block=[cute.size(tv_layout, mode=[0]), 1, 1],
)
def run_elementwise_apply_and_verify(
op,
M,
N,
dtype: Type[cutlass.Numeric],
skip_ref_check=False,
benchmark=True,
warmup_iterations=2,
iterations=100,
):
if not torch.cuda.is_available():
raise RuntimeError(f"Ampere GPU is required to run this example!")
print(f"\nRunning Elementwise Apply test with:")
print(f"Tensor dimensions: [{M}, {N}]")
print(f"Input and Output Data type: {dtype}")
print(f"Warmup iterations: {warmup_iterations}")
print(f"Measurement iterations: {iterations}\n")
torch_dtype = cutlass_torch.dtype(dtype)
# Allocate tensors with random values.
a = torch.randn(M, N, device=torch.device("cuda"), dtype=torch_dtype)
b = torch.randn(M, N, device=torch.device("cuda"), dtype=torch_dtype)
c = torch.zeros_like(a)
print(f"Input tensor shapes:")
print(f"a: {a.shape}, dtype: {a.dtype}")
print(f"b: {b.shape}, dtype: {b.dtype}")
print(f"c: {c.shape}, dtype: {c.dtype}\n")
epsilon = 1.2
if op in (operator.truediv, operator.floordiv):
b = torch.where(b == 0, torch.tensor(epsilon), b)
print("Compiling kernel with cute.compile ...")
start_time = time.time()
compilation_time = time.time() - start_time
print(f"Compilation time: {compilation_time:.4f} seconds")
print("Executing elementwise apply kernel...")
# Get current CUDA stream from PyTorch
torch_stream = torch.cuda.current_stream()
# Get the raw stream pointer as a CUstream
current_stream = cuda.CUstream(torch_stream.cuda_stream)
if not skip_ref_check:
elementwise_apply(
op, from_dlpack(a), from_dlpack(b), from_dlpack(c).mark_layout_dynamic()
)
print("Verifying results...")
torch.testing.assert_close(op(a, b), c)
print("Results verified successfully!")
if not benchmark:
return
# Create CUDA events for timing
start_event = cuda.cuEventCreate(cuda.CUevent_flags.CU_EVENT_DEFAULT)[1]
end_event = cuda.cuEventCreate(cuda.CUevent_flags.CU_EVENT_DEFAULT)[1]
# Warmup
for _ in range(warmup_iterations):
elementwise_apply(
op, from_dlpack(a), from_dlpack(b), from_dlpack(c).mark_layout_dynamic()
)
# Record start event
cuda.cuEventRecord(start_event, current_stream)
# Execute the kernel
for _ in range(iterations):
elementwise_apply(
op, from_dlpack(a), from_dlpack(b), from_dlpack(c).mark_layout_dynamic()
)
# Record end event
cuda.cuEventRecord(end_event, current_stream)
cuda.cuEventSynchronize(end_event)
# Calculate elapsed time
err, elapsed_time = cuda.cuEventElapsedTime(start_event, end_event)
avg_time = elapsed_time / iterations
# Print execution results
print(f"Kernel execution time: {avg_time:.4f} ms")
print(
f"Achieved memory throughput: {(3 * a.numel() * dtype.width // 8) / (avg_time / 1000) / 1e9:.2f} GB/s"
)
print(f"First few elements of result: \n{c[:3, :3]}")
# Destroy events
cuda.cuEventDestroy(start_event)
cuda.cuEventDestroy(end_event)
if __name__ == "__main__":
parser = argparse.ArgumentParser(
description="example of elementwise apply to demonstrate building elementwise kernels"
)
parser.add_argument("--M", default=128, type=int)
parser.add_argument("--N", default=128, type=int)
parser.add_argument("--op", default="add", type=str)
parser.add_argument("--warmup_iterations", default=2, type=int)
parser.add_argument("--iterations", default=100, type=int)
parser.add_argument("--skip_ref_check", action="store_true")
parser.add_argument("--benchmark", action="store_true")
args = parser.parse_args()
run_elementwise_apply_and_verify(
getattr(operator, args.op),
args.M,
args.N,
dtype=cutlass.Float32,
warmup_iterations=args.warmup_iterations,
iterations=args.iterations,
skip_ref_check=args.skip_ref_check,
benchmark=args.benchmark,
)
print("\nPASS")
File diff suppressed because it is too large Load Diff
+780
View File
@@ -0,0 +1,780 @@
# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: BSD-3-Clause
# Redistribution and use in source and binary forms, with or without
# modification, are permitted provided that the following conditions are met:
# 1. Redistributions of source code must retain the above copyright notice, this
# list of conditions and the following disclaimer.
# 2. Redistributions in binary form must reproduce the above copyright notice,
# this list of conditions and the following disclaimer in the documentation
# and/or other materials provided with the distribution.
# 3. Neither the name of the copyright holder nor the names of its
# contributors may be used to endorse or promote products derived from
# this software without specific prior written permission.
# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
import argparse
import time
from typing import Tuple
import cuda.bindings.driver as cuda
import torch
import cutlass
import cutlass.cute as cute
import cutlass.utils as utils
from cutlass.cute.runtime import from_dlpack
"""
A dense FP32 SIMT GEMM (C = A * B) example using CUTE DSL.
- Matrix A is MxK, A can be row-major("K") or column-major("M")
- Matrix B is NxK, B can be row-major("N") or column-major("K")
- Matrix C is MxN, C can be row-major("N") or column-major("M")
This GEMM kernel supports the following features:
- Utilizes FPU for matrix multiply-accumulate (MMA) operations
- Use multistage pipeline to overlap computation and memory access
* Shared memory pipeline: hides gmem-to-smem latency.
* Register pipeline: overlaps shared memory-to-register transfers with
computations and eliminates false data dependencies for
better parallelism.
- Use vectorized copies
- Add padding to reduce bank conflicts in global -> shared memory copies
- Use predication to avoid unnecessary copies or copies of stale data
This GEMM works as follows:
1. Load A and B matrices from global memory (GMEM) to shared memory (SMEM) using asynchronous copies.
2. Perform matrix multiply-accumulate (MMA) operations using simple fused multiply-add atomics.
3. Store results from registers (RMEM) to global memory (GMEM).
To run this example:
.. code-block:: bash
python examples/ampere/sgemm.py \
--mnk 8192,8192,8192 \
--a_major m --b_major n --c_major n
To collect performance with NCU profiler:
.. code-block:: bash
ncu python examples/ampere/sgemm.py \
--mnk 8192,8192,8192 \
--a_major m --b_major n --c_major n \
--skip_ref_check --iterations 2
Constraints:
* Supported input, output, and accumulator data types: fp32
* Default tile shape is set to be 128x128x8
* The contiguous dimension of A/B/C tensors must be at least 16 bytes aligned
"""
class SGemm:
def __init__(
self,
cta_tiler: Tuple[int, int, int] = (128, 128, 8),
num_stages: int = 3,
num_threads: int = 256,
):
self._cta_tiler = cta_tiler
self._num_stages = num_stages
self._num_threads = num_threads
assert num_threads > 0, "needs at least one thread"
assert num_threads % 16 == 0, "multiples of 16 required for MMA thread layout"
self._bM, self._bN, self._bK = self._cta_tiler
assert self._bM % 16 == 0, "multiple of 16 required for tile dimension M"
assert self._bN % 16 == 0, "multiple of 16 required for tile dimension N"
assert self._num_stages >= 3, "num_stages must be greater than or equal to 3"
@cute.jit
def __call__(
self,
mA: cute.Tensor,
mB: cute.Tensor,
mC: cute.Tensor,
epilogue_op: cutlass.Constexpr = lambda x: x,
):
self.a_major_mode = utils.LayoutEnum.from_tensor(mA)
self.b_major_mode = utils.LayoutEnum.from_tensor(mB)
self.c_major_mode = utils.LayoutEnum.from_tensor(mC)
# ///////////////////////////////////////////////////////////////////////////////
# Create layouts for shared memory for A and B:
# - sA/sB is m/n-major to vectorized copies from shared
# memory to registers. This is because the MMA layouts
# for sA/sB are also m/n-major
# - When gA/gB is k-major, pad 4 elements to reduce bank conflicts
# ///////////////////////////////////////////////////////////////////////////////
padding_a = 4 if self.a_major_mode == utils.LayoutEnum.ROW_MAJOR else 0
padding_b = 4 if self.b_major_mode == utils.LayoutEnum.ROW_MAJOR else 0
sA_layout = cute.make_layout(
(self._bM, self._bK, self._num_stages),
stride=(1, (self._bM + padding_a), self._bK * (self._bM + padding_a)),
)
sB_layout = cute.make_layout(
(self._bN, self._bK, self._num_stages),
stride=(1, (self._bN + padding_b), self._bK * (self._bN + padding_b)),
)
smem_size = cute.size_in_bytes(mA.element_type, sA_layout) + cute.size_in_bytes(
mB.element_type, sB_layout
)
# ///////////////////////////////////////////////////////////////////////////////
# Create copy layouts that will be used for asynchronous
# global memory -> shared memory copies:
# - The majorness of tA/tB follows the majorness of gA/gB
# - For k-major, these layouts will copy values one-by-one from
# from global memory, without vectorizing
# - For m/n-major, it will vectorize to a 128bit copy for faster
# data transfer between global and shared memory, as long
# as the alignment of the tensor allows it. Otherwise, it
# defaults to a non-vectorized copy
# ///////////////////////////////////////////////////////////////////////////////
tA = cute.make_layout(
(self._num_threads // self._bK, self._bK), stride=(self._bK, 1)
)
tB = cute.make_layout(
(self._num_threads // self._bK, self._bK), stride=(self._bK, 1)
)
vA = cute.make_layout((1, 1))
vB = cute.make_layout((1, 1))
atom_async_copy_A = cute.make_copy_atom(
cute.nvgpu.cpasync.CopyG2SOp(),
mA.element_type,
num_bits_per_copy=mA.element_type.width,
)
atom_async_copy_B = cute.make_copy_atom(
cute.nvgpu.cpasync.CopyG2SOp(),
mA.element_type,
num_bits_per_copy=mB.element_type.width,
)
if self.a_major_mode == utils.LayoutEnum.COL_MAJOR:
num_vectorized = 4 if (mA.layout.max_alignment % 16 == 0) else 1
atom_async_copy_A = cute.make_copy_atom(
cute.nvgpu.cpasync.CopyG2SOp(),
mA.element_type,
num_bits_per_copy=mA.element_type.width * num_vectorized,
)
major_mode_size = self._bM // num_vectorized
tA = cute.make_layout(
(major_mode_size, self._num_threads // major_mode_size),
stride=(1, major_mode_size),
)
vA = cute.make_layout((num_vectorized, 1))
if self.b_major_mode == utils.LayoutEnum.COL_MAJOR:
num_vectorized = 4 if (mB.layout.max_alignment % 16 == 0) else 1
atom_async_copy_B = cute.make_copy_atom(
cute.nvgpu.cpasync.CopyG2SOp(),
mA.element_type,
num_bits_per_copy=mB.element_type.width * num_vectorized,
)
major_mode_size = self._bN // num_vectorized
tB = cute.make_layout(
(major_mode_size, self._num_threads // major_mode_size),
stride=(1, major_mode_size),
)
vB = cute.make_layout((num_vectorized, 1))
tiled_copy_A = cute.make_tiled_copy_tv(atom_async_copy_A, tA, vA)
tiled_copy_B = cute.make_tiled_copy_tv(atom_async_copy_B, tB, vB)
# ///////////////////////////////////////////////////////////////////////////////
# Create layouts for GEMM:
# We tile an MMA atom across a tensor. `atoms_layout` is the layout
# of atoms in the tiled MMA. (Because we use an `MmaUniversalOp`,
# which has a trivial 1x1x1 MMA trait, `atoms_layout` is also
# simply the thread layout for C.) `permutation_tiler` reorders the
# elements of the tensor that the tiled MMA is applied to.
# Different combinations of `atoms_layout` and `permutation_tiler`
# values can create different MMA thread-value patterns.
#
# Here, the MMA layout is set so that each thread copies four
# consecutive elements from shared memory to registers.
# `permutation_tiler_M/N` maps the elements handled by each thread
# to the permuted element in the tensor.
# For increasing indices in the tensor, the thread ID that reads it is:
# - (without permutation) ==>
# 0 1 2 ... 15 0 1 2 ... 15 0 1 2 ... 15 0 1 2 ... 15 ......
# - (with permutation) ==>
# 0 0 0 0 1 1 1 1 2 2 2 2 ... 15 15 15 15 0 0 0 0 1 1 1 1 ......
# ///////////////////////////////////////////////////////////////////////////////
atoms_layout = cute.make_layout(
(self._num_threads // 16, 16, 1), stride=(16, 1, 0)
)
if self.c_major_mode == utils.LayoutEnum.COL_MAJOR:
atoms_layout = cute.make_layout(
(16, self._num_threads // 16, 1), stride=(1, 16, 0)
)
op = cute.nvgpu.MmaUniversalOp(cutlass.Float32)
permutation_tiler_M = cute.make_layout(
(atoms_layout.shape[0], 4), stride=(4, 1)
)
permutation_tiler_N = cute.make_layout(
(atoms_layout.shape[1], 4), stride=(4, 1)
)
tiled_mma = cute.make_tiled_mma(
op,
atoms_layout,
permutation_mnk=(permutation_tiler_M, permutation_tiler_N, None),
)
# grid_dim: ((m + BLK_M - 1) // BLK_M, (n + BLK_N - 1) // BLK_N, 1)
grid_dim = *cute.ceil_div(mC.shape, (self._bM, self._bN)), 1
self.kernel(
mA,
mB,
mC,
sA_layout,
sB_layout,
tiled_copy_A,
tiled_copy_B,
tiled_mma,
epilogue_op,
).launch(
grid=grid_dim,
block=[cute.size(atoms_layout), 1, 1],
smem=smem_size,
)
@cute.kernel
def kernel(
self,
mA: cute.Tensor,
mB: cute.Tensor,
mC: cute.Tensor,
sA_layout: cute.Layout,
sB_layout: cute.Layout,
tiled_copy_A: cute.TiledCopy,
tiled_copy_B: cute.TiledCopy,
tiled_mma: cute.TiledMma,
epilogue_op: cutlass.Constexpr = lambda x: x,
):
# Thread and block indices
tidx, tidy, tidz = cute.arch.thread_idx()
bidx, bidy, bidz = cute.arch.block_idx()
tiler_coord = (bidx, bidy, None)
thr_mma = tiled_mma.get_slice(tidx)
# ///////////////////////////////////////////////////////////////////////////////
# Get the appropriate tiles for this thread block.
# gA: (BLK_M, BLK_K, k), gB: (BLK_N, BLK_K, k), gC: (BLK_M, BLK_N)
# ///////////////////////////////////////////////////////////////////////////////
gA = cute.local_tile(
mA, tiler=self._cta_tiler, coord=tiler_coord, proj=(1, None, 1)
)
gB = cute.local_tile(
mB, tiler=self._cta_tiler, coord=tiler_coord, proj=(None, 1, 1)
)
gC = cute.local_tile(
mC, tiler=self._cta_tiler, coord=tiler_coord, proj=(1, 1, None)
)
# Move the pointer of gA/gB in the `-k`` direction, making the first
# tile (instead of the last one) irregular in shape when k is irregular.
# We first handle the irregular tile to avoid checking for this
# condition within the mainloop.
residue_k = mA.shape[1] - cutlass.Int32(self._bK) * gA.shape[2]
gA = cute.domain_offset((0, residue_k, 0), gA)
gB = cute.domain_offset((0, residue_k, 0), gB)
# ///////////////////////////////////////////////////////////////////////////////
# Get the appropriate tiles for this thread.
# sA: (BLK_M, BLK_K, PIPE) , sB: (BLK_N, BLK_K, PIPE)
# tAgA: (CPY, CPY_M, CPY_K, k) , tBgB: (CPY, CPY_N, CPY_K, k)
# tAsA: (CPY, CPY_M, CPY_K, PIPE) , tBsB: (CPY, CPY_N, CPY_K, PIPE)
# ///////////////////////////////////////////////////////////////////////////////
# Create shared memory buffer
smem = cutlass.utils.SmemAllocator()
sA = smem.allocate_tensor(mA.element_type, sA_layout, 16)
sB = smem.allocate_tensor(mB.element_type, sB_layout, 16)
thr_copy_A = tiled_copy_A.get_slice(tidx)
thr_copy_B = tiled_copy_B.get_slice(tidx)
tAgA = thr_copy_A.partition_S(gA)
tAsA = thr_copy_A.partition_D(sA)
tBgB = thr_copy_B.partition_S(gB)
tBsB = thr_copy_B.partition_D(sB)
# ///////////////////////////////////////////////////////////////////////////////
# Predicate: Mark indices that need to copy when the problem shape
# isn't a multiple of the tile shape. If tApA/B[i] is 0, then do not
# do the copy atom associated with index i.
# cA: (BLK_M, BLK_K) => (blk_m, blk_k)
# cB: (BLK_N, BLK_K) => (blk_n, blk_k)
# tAcA: (CPY, CPY_M, CPY_K) => (blk_m, blk_k)
# tBcB: (CPY, CPY_N, CPY_K) => (blk_n, blk_k)
# tApA: (rest_v, CPY_M, CPY_K), stride=(..., ..., 0)
# tBpB: (rest_v, CPY_N, CPY_K), stride=(..., ..., 0)
# CPY = (atom_v, rest_v)
# ///////////////////////////////////////////////////////////////////////////////
# Construct identity layout for sA and sB, used for predication
mcA = cute.make_identity_tensor(mA.shape)
mcB = cute.make_identity_tensor(mB.shape)
cA = cute.local_tile(
mcA, tiler=self._cta_tiler, coord=tiler_coord, proj=(1, None, 1)
)
cB = cute.local_tile(
mcB, tiler=self._cta_tiler, coord=tiler_coord, proj=(None, 1, 1)
)
cA = cute.domain_offset((0, residue_k, 0), cA)
cB = cute.domain_offset((0, residue_k, 0), cB)
# Repeat the partitioning with identity layouts
tAcA = thr_copy_A.partition_S(cA)
tBcB = thr_copy_B.partition_S(cB)
# Allocate predicate tensors for m and n
tApA = cute.make_fragment(
cute.make_layout(
(
tAsA.shape[0][1],
cute.size(tAsA, mode=[1]),
cute.size(tAsA, mode=[2]),
),
stride=(cute.size(tAsA, mode=[1]), 1, 0),
),
cutlass.Boolean,
)
tBpB = cute.make_fragment(
cute.make_layout(
(
tBsB.shape[0][1],
cute.size(tBsB, mode=[1]),
cute.size(tBsB, mode=[2]),
),
stride=(cute.size(tBsB, mode=[1]), 1, 0),
),
cutlass.Boolean,
)
# Allocate predicate tensors for m, n and k for residue k-tile
tApA_residue_k = cute.make_fragment(
cute.make_layout(
(
tAsA.shape[0][1],
cute.size(tAsA, mode=[1]),
cute.size(tAsA, mode=[2]),
),
stride=(
cute.size(tAsA, mode=[1]) * cute.size(tAsA, mode=[2]),
cute.size(tAsA, mode=[2]),
1,
),
),
cutlass.Boolean,
)
tBpB_residue_k = cute.make_fragment(
cute.make_layout(
(
tBsB.shape[0][1],
cute.size(tBsB, mode=[1]),
cute.size(tBsB, mode=[2]),
),
stride=(
cute.size(tBsB, mode=[1]) * cute.size(tBsB, mode=[2]),
cute.size(tBsB, mode=[2]),
1,
),
),
cutlass.Boolean,
)
# Set predicates for m/n bounds for mainloop
for rest_v in range(tApA.shape[0]):
for m in range(tApA.shape[1]):
tApA[rest_v, m, 0] = cute.elem_less(
tAcA[(0, rest_v), m, 0, 0][0], mA.shape[0]
)
for rest_v in range(tBpB.shape[0]):
for n in range(tBpB.shape[1]):
tBpB[rest_v, n, 0] = cute.elem_less(
tBcB[(0, rest_v), n, 0, 0][0], mB.shape[0]
)
# Set predicates for m/n/k bounds for residue k tile
for rest_v in range(tApA_residue_k.shape[0]):
for m in range(tApA_residue_k.shape[1]):
for k in range(tApA_residue_k.shape[2]):
coord_A = tAcA[(0, rest_v), m, k, 0]
tApA_residue_k[rest_v, m, k] = cute.elem_less(
(coord_A[0], cutlass.Int32(-1)), (mA.shape[0], coord_A[1])
)
for rest_v in range(tBpB_residue_k.shape[0]):
for n in range(tBpB_residue_k.shape[1]):
for k in range(tBpB_residue_k.shape[2]):
coord_B = tBcB[(0, rest_v), n, k, 0]
tBpB_residue_k[rest_v, n, k] = cute.elem_less(
(coord_B[0], cutlass.Int32(-1)), (mB.shape[0], coord_B[1])
)
# ///////////////////////////////////////////////////////////////////////////////
# Prefetch Prologue
# ///////////////////////////////////////////////////////////////////////////////
# Start async loads for 0th k-tile, where we take care of the k-residue
k_pipe_max = cute.size(tAsA, mode=[3])
k_tile_count = cute.size(tAgA, mode=[3])
gmem_pipe_read = cutlass.Int32(0)
cute.copy(
tiled_copy_A,
tAgA[None, None, None, gmem_pipe_read],
tAsA[None, None, None, 0],
pred=tApA_residue_k,
)
cute.copy(
tiled_copy_B,
tBgB[None, None, None, gmem_pipe_read],
tBsB[None, None, None, 0],
pred=tBpB_residue_k,
)
cute.arch.cp_async_commit_group()
gmem_pipe_read = (
gmem_pipe_read + 1
if gmem_pipe_read + 1 < k_tile_count
else cutlass.Int32(0)
)
# Start async loads for 1st k-tile onwards, no k-residue handling needed
for k_tile in range(1, k_pipe_max - 1):
if k_tile < k_tile_count:
cute.copy(
tiled_copy_A,
tAgA[None, None, None, gmem_pipe_read],
tAsA[None, None, None, k_tile],
pred=tApA,
)
cute.copy(
tiled_copy_B,
tBgB[None, None, None, gmem_pipe_read],
tBsB[None, None, None, k_tile],
pred=tBpB,
)
gmem_pipe_read = (
gmem_pipe_read + 1
if gmem_pipe_read + 1 < k_tile_count
else cutlass.Int32(0)
)
cute.arch.cp_async_commit_group()
# all tiles have been copied from global memory, so clear the
# predicate tensor
if k_tile_count < k_pipe_max:
for rest_v in range(tApA.shape[0]):
for m in range(tApA.shape[1]):
tApA[rest_v, m, 0] = cutlass.Boolean(0)
for rest_v in range(tBpB.shape[0]):
for n in range(tBpB.shape[1]):
tBpB[rest_v, n, 0] = cutlass.Boolean(0)
# ///////////////////////////////////////////////////////////////////////////////
# Define A/B partitioning and C accumulators.
# ///////////////////////////////////////////////////////////////////////////////
tCsA = thr_mma.partition_A(sA)
tCsB = thr_mma.partition_B(sB)
tCgC = thr_mma.partition_C(gC)
tCrA = tiled_mma.make_fragment_A(tCsA[None, None, None, 0])
tCrB = tiled_mma.make_fragment_B(tCsB[None, None, None, 0])
tCrC = tiled_mma.make_fragment_C(tCgC)
# Clear the accumulator
tCrC.fill(0.0)
# Current pipe index in smem to read from / write to
smem_pipe_read = cutlass.Int32(0)
smem_pipe_write = cutlass.Int32(k_pipe_max - 1)
tCsA_p = tCsA[None, None, None, smem_pipe_read]
tCsB_p = tCsB[None, None, None, smem_pipe_read]
# ///////////////////////////////////////////////////////////////////////////////
# PREFETCH register pipeline
# ///////////////////////////////////////////////////////////////////////////////
k_block_max = cute.size(tCrA, mode=[2])
if k_block_max > 1:
# Wait until our first prefetched tile is loaded in
cute.arch.cp_async_wait_group(k_pipe_max - 2)
cute.arch.barrier()
# Prefetch the first rmem from the first k-tile
cute.autovec_copy(tCsA_p[None, None, 0], tCrA[None, None, 0])
cute.autovec_copy(tCsB_p[None, None, 0], tCrB[None, None, 0])
# ///////////////////////////////////////////////////////////////////////////////
# Mainloop
# 1. Shared memory pipeline (gmem -> smem):
# The default smem pipeline depth is 3, meaning that for shared
# memory buffers, we allocate three times the size described by the
# CTA tiler. We prefetch 2 of these buffers before entering the main
# loop. Considering only the transfer from global memory to shared
# memory, the general structure of the mainloop is:
# (1) copy k-tile from gmem to smem;
# (2) perform gemm computation on k-tile;
# (3) wait for the next copy to finish.
# The `cute.arch.cp_async_wait_group(num_smem_stages - 2)` command
# waits for the number of unfinished 'copy' to be <= 1. The advantage
# of this approach is that it allows for simultaneous production
# (i.e., step (1)) and consumption (i.e., step (2)) of smem.
# A common misconception is to prefetch N buffers and rewrite
# the pipeline logic to wait on N-1 pending copies. The disadvantage
# of this approach is that it requires fully consuming a buffer in
# order to open an empty buffer for the next copy.
# 2. Register pipeline (smem -> register):
# Similarly, the register pipeline produces i+1, consumes i, and
# produces i+2... Notably, i and i+1 do not use the same register,
# eliminating dependencies on the same register for better parallelism.
# 3. Combining the smem and register pipelines results in the mainloop.
# ///////////////////////////////////////////////////////////////////////////////
for _ in cutlass.range_dynamic(k_tile_count, unroll=1):
for k_block in range(k_block_max):
if k_block == k_block_max - 1:
tCsA_p = tCsA[None, None, None, smem_pipe_read]
tCsB_p = tCsB[None, None, None, smem_pipe_read]
cute.arch.cp_async_wait_group(k_pipe_max - 2)
cute.arch.barrier()
# Load A, B from shared memory to registers for k_block + 1
k_block_next = (k_block + 1) % k_block_max # static
cute.autovec_copy(
tCsA_p[None, None, k_block_next],
tCrA[None, None, k_block_next],
)
cute.autovec_copy(
tCsB_p[None, None, k_block_next],
tCrB[None, None, k_block_next],
)
# Fetch next A: To better interleave global memory access and
# compute instructions, we intentionally use the sequence:
# copy A, perform GEMM, then copy B.
if k_block == 0:
cute.copy(
tiled_copy_A,
tAgA[None, None, None, gmem_pipe_read],
tAsA[None, None, None, smem_pipe_write],
# Use predicates because the m-mode may be irregular
pred=tApA,
)
# Thread-level register gemm for k_block
cute.gemm(
tiled_mma,
tCrC,
tCrA[None, None, k_block],
tCrB[None, None, k_block],
tCrC,
)
# Fetch next B and update smem pipeline read/write
if k_block == 0:
cute.copy(
tiled_copy_B,
tBgB[None, None, None, gmem_pipe_read],
tBsB[None, None, None, smem_pipe_write],
# Use predicates because the n-mode may be irregular
pred=tBpB,
)
cute.arch.cp_async_commit_group()
smem_pipe_write = smem_pipe_read
smem_pipe_read = smem_pipe_read + 1
if smem_pipe_read == k_pipe_max:
smem_pipe_read = cutlass.Int32(0)
# After copying all tiles, we avoid clearing the predicate
# tensor in the `mainloop` to prevent increasing its
# instruction count. Instead, we continue copying the
# first tile, though it won't be used. The 0-th tile is not
# copied due to its irregular shape, which could lead to
# illegal memory accesses.
gmem_pipe_read = (
gmem_pipe_read + 1
if gmem_pipe_read + 1 < k_tile_count
else cutlass.Int32(1)
)
# ///////////////////////////////////////////////////////////////////////////////
# Epilogue
# Applies the epilogue operation to the accumulated results and copies
# them without vectorization.
# ///////////////////////////////////////////////////////////////////////////////
cute.arch.cp_async_wait_group(0)
cute.arch.barrier()
tCrC.store(epilogue_op(tCrC.load()))
# predicate
cC = cute.make_identity_tensor(gC.shape)
tCpC = thr_mma.partition_C(cC)
predC = cute.make_fragment(tCrC.layout, cutlass.Boolean)
residue_m = mC.shape[0] - cutlass.Int32(self._bM) * bidx
residue_n = mC.shape[1] - cutlass.Int32(self._bN) * bidy
for i in range(cute.size(tCrC.shape)):
predC[i] = cute.elem_less(tCpC[i], (residue_m, residue_n))
numIterM = cute.size(tCrC, mode=[1])
numIterN = cute.size(tCrC, mode=[2])
atom = cute.make_copy_atom(cute.nvgpu.CopyUniversalOp(), mC.element_type)
cute.copy(atom, tCrC, tCgC, pred=predC)
return
def main(
a_major: str,
b_major: str,
c_major: str,
problem_shape: Tuple[int, int, int],
warmup_iterations: int = 2,
iterations: int = 100,
skip_ref_check: bool = False,
):
torch.manual_seed(1024)
M, N, K = problem_shape
# Create and permute tensor A/B/C
def create_and_permute_tensor(mode0, mode1, is_mode0_major, dtype):
# is_mode0_major: (mode1, mode0) -> (mode0, mode1)
# else: (mode0, mode1) -> (mode0, mode1)
shape = (mode1, mode0) if is_mode0_major else (mode0, mode1)
permute_order = (1, 0) if is_mode0_major else (0, 1)
return (
torch.empty(*shape, dtype=torch.int32)
.random_(-5, 5)
.to(dtype=dtype)
.permute(permute_order)
.cuda()
)
a = create_and_permute_tensor(M, K, a_major == "m", torch.float32)
b = create_and_permute_tensor(N, K, b_major == "n", torch.float32)
c = create_and_permute_tensor(M, N, c_major == "m", torch.float32)
divisibility_a = a.shape[1] if a_major == "k" else a.shape[0]
divisibility_b = b.shape[1] if b_major == "k" else b.shape[0]
divisibility_c = c.shape[1] if c_major == "n" else c.shape[0]
a_tensor = (
from_dlpack(a, assumed_align=16)
.mark_layout_dynamic(leading_dim=(1 if a_major == "k" else 0))
.mark_compact_shape_dynamic(
mode=(1 if a_major == "k" else 0),
divisibility=divisibility_a,
)
)
b_tensor = (
from_dlpack(b, assumed_align=16)
.mark_layout_dynamic(leading_dim=(1 if b_major == "k" else 0))
.mark_compact_shape_dynamic(
mode=(1 if b_major == "k" else 0),
divisibility=divisibility_b,
)
)
c_tensor = (
from_dlpack(c, assumed_align=16)
.mark_layout_dynamic(leading_dim=(1 if c_major == "n" else 0))
.mark_compact_shape_dynamic(
mode=(1 if c_major == "n" else 0),
divisibility=divisibility_c,
)
)
sgemm = SGemm()
print("Compiling kernel with cute.compile ...")
start_time = time.time()
gemm = cute.compile(sgemm, a_tensor, b_tensor, c_tensor)
compilation_time = time.time() - start_time
print(f"Compilation time: {compilation_time:.4f} seconds")
print("Executing GEMM kernel...")
# Get current CUDA stream from PyTorch
torch_stream = torch.cuda.current_stream()
# Get the raw stream pointer as a CUstream
current_stream = cuda.CUstream(torch_stream.cuda_stream)
# Create CUDA events for timing
start_event = cuda.cuEventCreate(cuda.CUevent_flags.CU_EVENT_DEFAULT)[1]
end_event = cuda.cuEventCreate(cuda.CUevent_flags.CU_EVENT_DEFAULT)[1]
# Warmup
for _ in range(warmup_iterations):
gemm(a_tensor, b_tensor, c_tensor)
# Use the current stream for CUDA events instead of the default stream
# Record start event
cuda.cuEventRecord(start_event, current_stream)
# Execute the kernel
for _ in range(iterations):
gemm(a_tensor, b_tensor, c_tensor)
# Record end event
cuda.cuEventRecord(end_event, current_stream)
cuda.cuEventSynchronize(end_event)
# Calculate elapsed time
err, elapsed_time = cuda.cuEventElapsedTime(start_event, end_event)
# Print execution results
print(f"Kernel execution time: {elapsed_time / iterations:.4f} ms")
# Destroy events
cuda.cuEventDestroy(start_event)
cuda.cuEventDestroy(end_event)
if not skip_ref_check:
print("Verifying results...")
ref = torch.einsum("mk,nk->mn", a, b)
torch.testing.assert_close(c.cpu(), ref.cpu(), atol=1e-03, rtol=1e-05)
print("Results verified successfully!")
if __name__ == "__main__":
def parse_comma_separated_ints(s: str) -> Tuple[int, ...]:
try:
return tuple(int(x.strip()) for x in s.split(","))
except ValueError:
raise argparse.ArgumentTypeError(
"Invalid format. Expected comma-separated integers."
)
parser = argparse.ArgumentParser()
parser.add_argument(
"--mnk", type=parse_comma_separated_ints, default=(256, 256, 64)
)
parser.add_argument("--a_major", choices=["k", "m"], default="k")
parser.add_argument("--b_major", choices=["k", "n"], default="k")
parser.add_argument("--c_major", choices=["n", "m"], default="n")
parser.add_argument("--warmup_iterations", default=2, type=int)
parser.add_argument("--iterations", default=100, type=int)
parser.add_argument("--skip_ref_check", action="store_true")
args = parser.parse_args()
print("Running SIMT GEMM example:")
main(
args.a_major,
args.b_major,
args.c_major,
args.mnk,
args.warmup_iterations,
args.iterations,
args.skip_ref_check,
)
print("PASS")
@@ -0,0 +1,968 @@
# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: BSD-3-Clause
# Redistribution and use in source and binary forms, with or without
# modification, are permitted provided that the following conditions are met:
# 1. Redistributions of source code must retain the above copyright notice, this
# list of conditions and the following disclaimer.
# 2. Redistributions in binary form must reproduce the above copyright notice,
# this list of conditions and the following disclaimer in the documentation
# and/or other materials provided with the distribution.
# 3. Neither the name of the copyright holder nor the names of its
# contributors may be used to endorse or promote products derived from
# this software without specific prior written permission.
# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
import argparse
import math
import time
from typing import Tuple, Type
import cuda.bindings.driver as cuda
import torch
import cutlass
import cutlass.cute as cute
import cutlass.torch as cutlass_torch
import cutlass.utils as utils
from cutlass.cute.runtime import from_dlpack
"""
A dense GEMM (C = A * B) example for the NVIDIA Ampere architecture using CUTE DSL.
- Matrix A is MxKxL, L is batch dimension, A can be row-major("K") or column-major("M")
- Matrix B is NxKxL, L is batch dimension, B can be row-major("N") or column-major("K")
- Matrix C is MxNxL, L is batch dimension, C can be row-major("N") or column-major("M")
This GEMM kernel supports the following features:
- Utilizes Ampere's tensor cores for matrix multiply-accumulate (MMA) operations
- Supports multi-stage pipeline to overlap computation and memory access
- Implements shared memory buffering for epilogue to increase coalesed global memory access
This GEMM works as follows:
1. Load A and B matrices from global memory (GMEM) to shared memory (SMEM) using asynchronous copies.
2. Perform matrix multiply-accumulate (MMA) operations.
3. Store results from registers (RMEM) to shared memory (SMEM), then to global memory (GMEM).
The Ampere tensor core instruction used operates as follows:
- Read matrix A from SMEM
- Read matrix B from SMEM
- Perform MMA operation and store the result in Accumulator(register)
To run this example:
.. code-block:: bash
python examples/ampere/tensorop_gemm.py \
--mnkl 8192,8192,8192,1 --atom_layout_mnk 2,2,1 \
--ab_dtype Float16 \
--c_dtype Float16 --acc_dtype Float32 \
--a_major m --b_major n --c_major n
The above example command computes with M=8192, N=8192, K=8192,
batch_count=1. The atom layout's shape is 2x2x1 and the input, mma
accumulator, and output data type are set as fp16, fp32 and fp16,
respectively.
To collect performance with NCU profiler:
.. code-block:: bash
ncu python examples/ampere/tensorop_gemm.py \
--mnkl 8192,8192,8192,1 --atom_layout_mnk 2,2,1 \
--ab_dtype Float16 \
--c_dtype Float16 --acc_dtype Float32 \
--a_major m --b_major n --c_major n \
--skip_ref_check --iterations 2
Constraints:
* Supported input and output data types: fp16
* Support accumulator data types: f32
* Default tile shape is set to be 128x128x32
* Atom layout's MNK shape is set so that tile shape can be divided by MMA
instruction shape
* The contiguous dimension of A/B/C tensors must be at least 16 bytes aligned,
i.e, number of elements is a multiple of 8
"""
class TensorOpGemm:
def __init__(
self,
ab_dtype: Type[cutlass.Numeric],
c_dtype: Type[cutlass.Numeric],
acc_dtype: Type[cutlass.Numeric],
atom_layout_mnk: Tuple[int, int, int],
):
self.ab_dtype = ab_dtype
self.c_dtype = c_dtype
self.acc_dtype = acc_dtype
self.cta_tiler = (128, 128, 32)
self.num_stages = 3
self.atom_layout_mnk = atom_layout_mnk
atom_lay_M, atom_lay_N, atom_lay_K = self.atom_layout_mnk
self.num_threads = atom_lay_M * atom_lay_N * atom_lay_K * 32
self.bM, self.bN, self.bK = self.cta_tiler
self.mma_inst_shape = (16, 8, 16)
mmaM, mmaN, mmaK = self.mma_inst_shape
assert (
self.bM % (atom_lay_M * mmaM) == 0
), "bM must be divisible by MMA instruction"
assert (
self.bN % (atom_lay_N * mmaN) == 0
), "bN must be divisible by MMA instruction"
assert atom_lay_K == 1, "this example does not support atom layout K > 1"
assert self.bK % mmaK == 0, "bK must be divisible by MMA instruction"
assert self.num_stages >= 3, "num_stages must be greater than or equal to 3"
@cute.jit
def __call__(
self,
mA: cute.Tensor,
mB: cute.Tensor,
mC: cute.Tensor,
epilogue_op: cutlass.Constexpr = lambda x: x,
):
# The grid divides the problems's M, N, and L dimensions by the
# respective modes of the tile shape (bM, bN, 1). The K dimension is
# handled within a block via a multistage process.
self.a_major_mode = utils.LayoutEnum.from_tensor(mA)
self.b_major_mode = utils.LayoutEnum.from_tensor(mB)
self.c_major_mode = utils.LayoutEnum.from_tensor(mC)
# ///////////////////////////////////////////////////////////////////////////////
# Shared memory layout:
# ///////////////////////////////////////////////////////////////////////////////
# Creates a layout with the size required for the provided tile
# size and num stages (stages are used for K dimension) that is also
# sectioned into 64x8 or 8x32 layout atoms. The swizzle is set so that
# the atom for the shared memory -> register copy does not encounter
# bank conflicts
# assume the input is 16B align
ab_copy_bits = 128
sA_layout = self._make_smem_layout_AB(
mA.element_type,
self.a_major_mode,
ab_copy_bits,
(self.cta_tiler[0], self.cta_tiler[2], self.num_stages),
)
sB_layout = self._make_smem_layout_AB(
mB.element_type,
self.b_major_mode,
ab_copy_bits,
(self.cta_tiler[1], self.cta_tiler[2], self.num_stages),
)
# Creates a similar layout but without num_stages or layout atoms
sC_layout = self._make_smem_layout_C(
mC.element_type,
self.c_major_mode,
ab_copy_bits,
(self.cta_tiler[0], self.cta_tiler[1]),
)
# Shared memory allocated for operations with A, B will be
# overwritten for operations on C. This is to improve performance
# by reducing the size of shared memory requested by each block
smem_size = max(
cute.size_in_bytes(mC.element_type, sC_layout),
cute.size_in_bytes(mA.element_type, sA_layout)
+ cute.size_in_bytes(mB.element_type, sB_layout),
)
# ///////////////////////////////////////////////////////////////////////////////
# Tiled copy:
# The majorness of tA/tB/tC follows the majorness of gA/gB/gC,
# enabling merged accesses to global memory for faster data
# transfer between global and shared memory.
# ///////////////////////////////////////////////////////////////////////////////
# Create a copy atom for a global to shared memory asynchronous copy
atom_async_copy = cute.make_copy_atom(
cute.nvgpu.cpasync.CopyG2SOp(
cache_mode=cute.nvgpu.cpasync.LoadCacheMode.GLOBAL
),
mA.element_type,
num_bits_per_copy=ab_copy_bits,
)
# Create thread layouts for tiled copy from the copy atom where the
# thread layout simply follows the leading dimension of the tensor
tiled_copy_A = self._make_gmem_tiled_copy_AB(
atom_async_copy, mA.element_type, self.a_major_mode, ab_copy_bits
)
tiled_copy_B = self._make_gmem_tiled_copy_AB(
atom_async_copy, mB.element_type, self.b_major_mode, ab_copy_bits
)
# Creates a synchonous copy atom and thread layouts for the epilogue
c_copy_bits = 128
atom_sync_copy = cute.make_copy_atom(
cute.nvgpu.CopyUniversalOp(),
mC.element_type,
num_bits_per_copy=c_copy_bits,
)
tiled_copy_C = self._make_gmem_tiled_copy_C(
atom_sync_copy, mC.element_type, self.c_major_mode, c_copy_bits
)
# ///////////////////////////////////////////////////////////////////////////////
# Tiled MMA
# ///////////////////////////////////////////////////////////////////////////////
# Creates a mma atom with 16x8x16 shape for MNK
op = cute.nvgpu.warp.MmaF16BF16Op(
self.ab_dtype, self.acc_dtype, self.mma_inst_shape
)
permutation_mnk = (
self.atom_layout_mnk[0] * self.mma_inst_shape[0],
# if atom layout's N-mode is 1, to leverage the largest coalesced
# shared memory -> register copy, set the tiled mma's N mode to 16
self.atom_layout_mnk[1] * self.mma_inst_shape[1] * 2,
self.atom_layout_mnk[2] * self.mma_inst_shape[2],
)
# Created a tiled mma that tiles the atom according to specified layout.
# For a 2x2x1 atom layout, the mma atom is duplicated 4 times, twice
# across M and twice across N
tC = cute.make_layout(self.atom_layout_mnk)
tiled_mma = cute.make_tiled_mma(
op,
tC,
permutation_mnk=permutation_mnk,
)
# grid_dim: ((m + BLK_M - 1) // BLK_M, (n + BLK_N - 1) // BLK_N, l)
grid_dim = cute.ceil_div(mC.shape, (self.bM, self.bN, 1))
self.kernel(
mA,
mB,
mC,
sA_layout,
sB_layout,
sC_layout,
tiled_copy_A,
tiled_copy_B,
tiled_copy_C,
tiled_mma,
epilogue_op,
).launch(
grid=grid_dim,
block=[self.num_threads, 1, 1],
smem=smem_size,
)
@cute.kernel
def kernel(
self,
mA: cute.Tensor,
mB: cute.Tensor,
mC: cute.Tensor,
sA_layout: cute.ComposedLayout,
sB_layout: cute.ComposedLayout,
sC_layout: cute.ComposedLayout,
tiled_copy_A: cute.TiledCopy,
tiled_copy_B: cute.TiledCopy,
tiled_copy_C: cute.TiledCopy,
tiled_mma: cute.TiledMma,
epilogue_op: cutlass.Constexpr = lambda x: x,
):
# Thread index, block index
tidx, _, _ = cute.arch.thread_idx()
bidx, bidy, bidz = cute.arch.block_idx()
tiler_coord = (bidx, bidy, None)
# ///////////////////////////////////////////////////////////////////////////////
# Get the appropriate tiles for this thread block.
# gA: (BLK_M, BLK_N, k), gB: (BLK_N, BLK_K, k), gC: (BLK_M, BLK_N)
# ///////////////////////////////////////////////////////////////////////////////
gA = cute.local_tile(
mA[None, None, bidz],
tiler=self.cta_tiler,
coord=tiler_coord,
proj=(1, None, 1),
)
gB = cute.local_tile(
mB[None, None, bidz],
tiler=self.cta_tiler,
coord=tiler_coord,
proj=(None, 1, 1),
)
gC = cute.local_tile(
mC[None, None, bidz],
tiler=self.cta_tiler,
coord=tiler_coord,
proj=(1, 1, None),
)
# By default, if the tensor k mode does not divide into the tile k
# size, then last tiles in the k dimension are irregular.
# Instead, make the first tiles irregular when k is irregular.
# This allows us to handle the irregular tile first to avoid
# checking for this condition within the mainloop.
# residual_k is a negative number indicating the amount needed to
# shift the pointer by in dimension k
residual_k = cute.size(mA, mode=[1]) - cutlass.Int32(self.bK) * cute.size(
gA, mode=[2]
)
# move the pointer of gA/gB in the `-k` direction
gA = cute.domain_offset((0, residual_k, 0), gA)
gB = cute.domain_offset((0, residual_k, 0), gB)
# input is 16B aligned
gA = cute.make_tensor(gA.iterator.align(16), gA.layout)
gB = cute.make_tensor(gB.iterator.align(16), gB.layout)
# Construct identity layout for sA and sB (mirrors global tensors,
# used for predication only)
mcA = cute.make_identity_tensor(mA.layout.shape)
mcB = cute.make_identity_tensor(mB.layout.shape)
cA = cute.local_tile(
mcA[None, None, bidz],
tiler=self.cta_tiler,
coord=tiler_coord,
proj=(1, None, 1),
)
cB = cute.local_tile(
mcB[None, None, bidz],
tiler=self.cta_tiler,
coord=tiler_coord,
proj=(None, 1, 1),
)
cA = cute.domain_offset((0, residual_k, 0), cA)
cB = cute.domain_offset((0, residual_k, 0), cB)
# ///////////////////////////////////////////////////////////////////////////////
# Create shared memory buffers and get the appropriate fragments for this thread.
# sA: (BLK_M, BLK_K, PIPE) , sB: (BLK_N, BLK_K, PIPE)
# tAgA: (CPY, CPY_M, CPY_K, k) , tBgB: (CPY, CPY_N, CPY_K, k)
# tAsA: (CPY, CPY_M, CPY_K, PIPE) , tBsB: (CPY, CPY_N, CPY_K, PIPE)
# ///////////////////////////////////////////////////////////////////////////////
# Shared memory buffer
smem = cutlass.utils.SmemAllocator()
sA = smem.allocate_tensor(mA.element_type, sA_layout, 16)
sB = smem.allocate_tensor(mB.element_type, sB_layout, 16)
sC = cute.make_tensor(
cute.recast_ptr(sA.iterator, dtype=self.c_dtype), sC_layout
)
thr_copy_A = tiled_copy_A.get_slice(tidx)
thr_copy_B = tiled_copy_B.get_slice(tidx)
thr_copy_C = tiled_copy_C.get_slice(tidx)
tAgA = thr_copy_A.partition_S(gA)
tAsA = thr_copy_A.partition_D(sA)
tBgB = thr_copy_B.partition_S(gB)
tBsB = thr_copy_B.partition_D(sB)
tCsC_epilogue = thr_copy_C.partition_S(sC)
tCgC_epilogue = thr_copy_C.partition_D(gC)
# Repeat the partitioning with identity layouts
tAcA = thr_copy_A.partition_S(cA)
tBcB = thr_copy_B.partition_S(cB)
# ///////////////////////////////////////////////////////////////////////////////
# Predicate: Mark indices that need to copy when problem_shape isn't a multiple
# of tile_shape
# ///////////////////////////////////////////////////////////////////////////////
# For predication over the tensors A (M/K), B (N/K), and (in the
# epilogue) C (M/N), we will compute it in a fashion similar to an
# outer product. The predication along one of the dimensions is
# evaluated and stored in a predication tensor. Then, the
# predication for the remaining dimension is handled later via an
# if/else branch at the copy.
# For A and B, predication booleans along M/N are stored in a
# predication tensor and along K is handled via a if/else branch.
# Allocate predicate tensors for M and N. Predication is checked
# at the granularity of a copy atom, so the predicate tensor does not
# need separate booleans for individual elements within a copy
# atom (for example, the elements of tAgA.shape[0][0].)
tApA = cute.make_fragment(
cute.make_layout(
(
tAgA.shape[0][1],
cute.size(tAgA, mode=[1]),
cute.size(tAgA, mode=[2]),
),
stride=(cute.size(tAgA, mode=[1]), 1, 0),
),
cutlass.Boolean,
)
tBpB = cute.make_fragment(
cute.make_layout(
(
tBsB.shape[0][1],
cute.size(tBsB, mode=[1]),
cute.size(tBsB, mode=[2]),
),
stride=(cute.size(tBsB, mode=[1]), 1, 0),
),
cutlass.Boolean,
)
# Set predicates for M/N bounds
for rest_v in range(tApA.shape[0]):
for m in range(tApA.shape[1]):
tApA[rest_v, m, 0] = cute.elem_less(
tAcA[(0, rest_v), m, 0, 0][0], mA.shape[0]
)
for rest_v in range(tBpB.shape[0]):
for n in range(tBpB.shape[1]):
tBpB[rest_v, n, 0] = cute.elem_less(
tBcB[(0, rest_v), n, 0, 0][0], mB.shape[0]
)
# ///////////////////////////////////////////////////////////////////////////////
# Prefetch Prologue
# ///////////////////////////////////////////////////////////////////////////////
# Clear the smem tiles to account for predicated off loads
tAsA.fill(0)
tBsB.fill(0)
cute.arch.sync_threads()
# Start async loads for the first k-tile. Here we take care of the k residue
# via if/else check along the k dimension. Because we shifted the identity tensor
# by the residue_k and because the identity tensor is a counting tensor, the
# values of any identity tensor element that is poison is less than -1
num_smem_stages = cute.size(tAsA, mode=[3])
k_tile_count = cute.size(tAgA, mode=[3])
k_tile_index = cutlass.Int32(0)
for k in range(tApA.shape[2]):
if cute.elem_less(cutlass.Int32(-1), tAcA[0, 0, k, 0][1]):
cute.copy(
tiled_copy_A,
tAgA[None, None, k, k_tile_index],
tAsA[None, None, k, 0],
pred=tApA[None, None, k],
)
for k in range(tBpB.shape[2]):
if cute.elem_less(cutlass.Int32(-1), tBcB[0, 0, k, 0][1]):
cute.copy(
tiled_copy_B,
tBgB[None, None, k, k_tile_index],
tBsB[None, None, k, 0],
pred=tBpB[None, None, k],
)
k_tile_index = k_tile_index + 1
cute.arch.cp_async_commit_group()
# Start async loads for rest of the k-tiles
for k_tile in range(1, num_smem_stages - 1):
if k_tile == k_tile_count:
tApA.fill(0)
tBpB.fill(0)
cute.copy(
tiled_copy_A,
tAgA[None, None, None, k_tile_index],
tAsA[None, None, None, k_tile],
pred=tApA,
)
cute.copy(
tiled_copy_B,
tBgB[None, None, None, k_tile_index],
tBsB[None, None, None, k_tile],
pred=tBpB,
)
k_tile_index = k_tile_index + 1
cute.arch.cp_async_commit_group()
# ///////////////////////////////////////////////////////////////////////////////
# Tile MMA compute thread partitions and allocate accumulators
# ///////////////////////////////////////////////////////////////////////////////
thr_mma = tiled_mma.get_slice(tidx)
tCsA = thr_mma.partition_A(sA)
tCsB = thr_mma.partition_B(sB)
tCsC = thr_mma.partition_C(sC)
tCgC = thr_mma.partition_C(gC)
tCrA = tiled_mma.make_fragment_A(tCsA[None, None, None, 0])
tCrB = tiled_mma.make_fragment_B(tCsB[None, None, None, 0])
tCrC = tiled_mma.make_fragment_C(tCgC)
# Clear the accumulator
tCrC.fill(0.0)
# ///////////////////////////////////////////////////////////////////////////////
# Copy Atom A/B retiling
# ///////////////////////////////////////////////////////////////////////////////
# Create the copy atoms for the copy from shared memory to register
atom_copy_s2r_A = cute.make_copy_atom(
cute.nvgpu.warp.LdMatrix8x8x16bOp(
self.a_major_mode != utils.LayoutEnum.ROW_MAJOR, 4
),
mA.element_type,
)
atom_copy_s2r_B = cute.make_copy_atom(
cute.nvgpu.warp.LdMatrix8x8x16bOp(
self.b_major_mode != utils.LayoutEnum.ROW_MAJOR, 4
),
mB.element_type,
)
# Creates the tiled copy so that it matches the thread-value layout
# expected by the tiled mma
tiled_copy_s2r_A = cute.make_tiled_copy(
atom_copy_s2r_A,
layout_tv=tiled_mma.tv_layout_A_tiled,
tiler_mn=(tiled_mma.get_tile_size(0), tiled_mma.get_tile_size(2)),
)
tiled_copy_s2r_B = cute.make_tiled_copy(
atom_copy_s2r_B,
layout_tv=tiled_mma.tv_layout_B_tiled,
tiler_mn=(tiled_mma.get_tile_size(1), tiled_mma.get_tile_size(2)),
)
thr_copy_ldmatrix_A = tiled_copy_s2r_A.get_slice(tidx)
thr_copy_ldmatrix_B = tiled_copy_s2r_B.get_slice(tidx)
tCsA_copy_view = thr_copy_ldmatrix_A.partition_S(sA)
tCrA_copy_view = thr_copy_ldmatrix_A.retile(tCrA)
tCsB_copy_view = thr_copy_ldmatrix_B.partition_S(sB)
tCrB_copy_view = thr_copy_ldmatrix_B.retile(tCrB)
# Current pipe index in smem to read from / write to
smem_pipe_read = 0
smem_pipe_write = num_smem_stages - 1
tCsA_p = tCsA_copy_view[None, None, None, smem_pipe_read]
tCsB_p = tCsB_copy_view[None, None, None, smem_pipe_read]
# ///////////////////////////////////////////////////////////////////////////////
# PREFETCH register pipeline
# ///////////////////////////////////////////////////////////////////////////////
num_k_block = cute.size(tCrA, mode=[2])
if num_k_block > 1:
# Wait until our first prefetched tile is loaded in
cute.arch.cp_async_wait_group(num_smem_stages - 2)
cute.arch.sync_threads()
# Prefetch the first k-block rmem from the first k-tile
cute.copy(
tiled_copy_s2r_A,
tCsA_p[None, None, 0],
tCrA_copy_view[None, None, 0],
)
cute.copy(
tiled_copy_s2r_B,
tCsB_p[None, None, 0],
tCrB_copy_view[None, None, 0],
)
# ///////////////////////////////////////////////////////////////////////////////
# Mainloop
# 1. Shared memory pipeline (gmem -> smem):
# The default smem pipeline depth is 3, meaning that for shared
# memory buffers, we allocate three times the size described by the
# CTA tiler. We prefetch 2 of these buffers before entering the main
# loop. Considering only the transfer from global memory to shared
# memory, the general structure of the mainloop is:
# (1) copy k-tile from gmem to smem;
# (2) perform gemm computation on k-tile;
# (3) wait for the next copy to finish.
# The `cute.arch.cp_async_wait_group(num_smem_stages - 2)` command
# waits for the number of unfinished 'copy' to be <= 1. The advantage
# of this approach is that it allows for simultaneous production
# (i.e., step (1)) and consumption (i.e., step (2)) of smem.
# A common misconception is to prefetch N buffers and rewrite
# the pipeline logic to wait on N-1 pending copies. The disadvantage
# of this approach is that it requires fully consuming a buffer in
# order to open an empty buffer for the next copy.
# 2. Register pipeline (smem -> register):
# Similarly, the register pipeline produces i+1, consumes i, and
# produces i+2... Notably, i and i+1 do not use the same register,
# eliminating dependencies on the same register for better parallelism.
# 3. Combining the smem and register pipelines results in the mainloop.
# ///////////////////////////////////////////////////////////////////////////////
for k_tile in cutlass.range_dynamic(k_tile_count, unroll=1):
for k_block in range(num_k_block):
if k_block == num_k_block - 1:
tCsA_p = tCsA_copy_view[None, None, None, smem_pipe_read]
tCsB_p = tCsB_copy_view[None, None, None, smem_pipe_read]
cute.arch.cp_async_wait_group(num_smem_stages - 2)
cute.arch.sync_threads()
# Load A, B from shared memory to registers for k_block + 1
k_block_next = (k_block + 1) % num_k_block # static
cute.copy(
tiled_copy_s2r_A,
tCsA_p[None, None, k_block_next],
tCrA_copy_view[None, None, k_block_next],
)
cute.copy(
tiled_copy_s2r_B,
tCsB_p[None, None, k_block_next],
tCrB_copy_view[None, None, k_block_next],
)
# Fetch next A: To better interleave global memory access and compute
# instructions, we intentionally use the sequence: copy A, perform GEMM,
# then copy B.
if k_block == 0:
if k_tile + num_smem_stages - 1 < k_tile_count:
cute.copy(
tiled_copy_A,
tAgA[None, None, None, k_tile_index],
tAsA[None, None, None, smem_pipe_write],
pred=tApA,
)
# Thread-level register gemm for k_block
cute.gemm(
tiled_mma,
tCrC,
tCrA[None, None, k_block],
tCrB[None, None, k_block],
tCrC,
)
# Fetch next B and update smem pipeline read/write
if k_block == 0:
if k_tile + num_smem_stages - 1 < k_tile_count:
cute.copy(
tiled_copy_B,
tBgB[None, None, None, k_tile_index],
tBsB[None, None, None, smem_pipe_write],
pred=tBpB,
)
k_tile_index = k_tile_index + 1
cute.arch.cp_async_commit_group()
smem_pipe_write = smem_pipe_read
smem_pipe_read = smem_pipe_read + 1
if smem_pipe_read == num_smem_stages:
smem_pipe_read = 0
# Sync before epilogue
cute.arch.cp_async_wait_group(0)
cute.arch.sync_threads()
# ///////////////////////////////////////////////////////////////////////////////
# Epilogue with fusion
# ///////////////////////////////////////////////////////////////////////////////
tCrD = cute.make_fragment_like(tCrC, self.c_dtype)
tCrD[None] = epilogue_op(tCrC.load()).to(self.c_dtype)
# Copy results of D back to shared memory
cute.autovec_copy(tCrD, tCsC)
# Create counting tensor for C
ceilM, ceilN, _ = cute.ceil_div(mC.shape, (self.bM, self.bN, 1))
mcC = cute.make_identity_tensor(
(
cute.size(ceilM) * self.cta_tiler[0],
cute.size(ceilN) * self.cta_tiler[1],
1,
)
)
cC = cute.local_tile(
mcC[None, None, bidz],
tiler=self.cta_tiler,
coord=tiler_coord,
proj=(1, 1, None),
)
tCcC = thr_copy_C.partition_S(cC)
tCrC_epilogue = cute.make_fragment_like(tCsC_epilogue)
# Wait for all writes to shared memory to finish before starting copies
# using the new layouts
cute.arch.sync_threads()
cute.autovec_copy(tCsC_epilogue, tCrC_epilogue)
# Create predication tensor for m
tCpC = cute.make_fragment(
cute.make_layout(
(
tCgC_epilogue.shape[0][1],
cute.size(tCgC_epilogue, mode=[1]),
cute.size(tCgC_epilogue, mode=[2]),
),
stride=(cute.size(tCgC_epilogue, mode=[1]), 1, 0),
),
cutlass.Boolean,
)
for rest_v in range(tCpC.shape[0]):
for m in range(tCpC.shape[1]):
tCpC[rest_v, m, 0] = cute.elem_less(
tCcC[(0, rest_v), m, 0][0], mC.shape[0]
)
# Copy to global memory using better vectorization
for rest_v in range(tCpC.shape[0]):
for n in range(tCpC.shape[2]):
if cute.elem_less(tCcC[(0, rest_v), 0, n][1], mC.shape[1]):
cute.copy(
tiled_copy_C,
tCrC_epilogue[None, None, n],
tCgC_epilogue[None, None, n],
pred=tCpC[None, None, n],
)
return
def _make_smem_layout_AB(self, dtype, major_mode, copy_bits, smem_tiler):
major_mode_size = (
smem_tiler[1] if major_mode == utils.LayoutEnum.ROW_MAJOR else smem_tiler[0]
)
major_mode_size = 64 if major_mode_size >= 64 else major_mode_size
swizzle_bits = int(math.log2(major_mode_size * dtype.width // copy_bits))
swizzle_bits = min(swizzle_bits, 3)
layout_atom_outer = (
cute.make_layout((8, major_mode_size), stride=(major_mode_size, 1))
if major_mode == utils.LayoutEnum.ROW_MAJOR
else cute.make_layout((major_mode_size, 8), stride=(1, major_mode_size))
)
layout_atom = cute.make_composed_layout(
cute.make_swizzle(swizzle_bits, 3, 3),
0,
layout_atom_outer,
)
layout = cute.tile_to_shape(layout_atom, smem_tiler, (0, 1, 2))
return layout
def _make_smem_layout_C(self, dtype, major_mode, copy_bits, smem_tiler):
major_mode_size = (
smem_tiler[1] if major_mode == utils.LayoutEnum.ROW_MAJOR else smem_tiler[0]
)
swizzle_bits = int(math.log2(major_mode_size * dtype.width // copy_bits))
swizzle_bits = min(swizzle_bits, 3)
layout_atom_outer = (
cute.make_layout((8, major_mode_size), stride=(major_mode_size, 1))
if major_mode == utils.LayoutEnum.ROW_MAJOR
else cute.make_layout((major_mode_size, 8), stride=(1, major_mode_size))
)
layout_atom = cute.make_composed_layout(
cute.make_swizzle(swizzle_bits, 3, 4),
0,
layout_atom_outer,
)
# Due to the thread layout of the mma, remove swizzle in C to
# prevent shared memory fragments owned by an single thread from
# holding swizzles
if major_mode == utils.LayoutEnum.COL_MAJOR:
layout_atom = cute.make_composed_layout(
cute.make_swizzle(0, 3, 4), 0, layout_atom_outer
)
layout = cute.tile_to_shape(
layout_atom,
smem_tiler,
(0, 1),
)
return layout
def _make_gmem_tiled_copy_AB(self, atom_copy, dtype, major_mode, copy_bits):
copy_elems = copy_bits // dtype.width
shape_dim_1 = cute.size(self.bK) // copy_elems
# thread layout for copy
thread_layout = cute.make_layout(
(self.num_threads // shape_dim_1, shape_dim_1), stride=(shape_dim_1, 1)
)
if major_mode != utils.LayoutEnum.ROW_MAJOR:
shape_dim_0 = cute.size(self.bM) // copy_elems
thread_layout = cute.make_layout(
(shape_dim_0, self.num_threads // shape_dim_0), stride=(1, shape_dim_0)
)
# Value layout for copy
value_layout = (
cute.make_layout((1, copy_elems))
if major_mode == utils.LayoutEnum.ROW_MAJOR
else cute.make_layout((copy_elems, 1))
)
return cute.make_tiled_copy_tv(atom_copy, thread_layout, value_layout)
def _make_gmem_tiled_copy_C(self, atom_copy, dtype, major_mode, copy_bits):
copy_elems = copy_bits // dtype.width
shape_dim_1 = cute.size(self.bN) // copy_elems
# thread layout for copy
thread_layout = cute.make_layout(
(self.num_threads // shape_dim_1, shape_dim_1), stride=(shape_dim_1, 1)
)
if major_mode != utils.LayoutEnum.ROW_MAJOR:
shape_dim_0 = cute.size(self.bM) // copy_elems
thread_layout = cute.make_layout(
(shape_dim_0, self.num_threads // shape_dim_0), stride=(1, shape_dim_0)
)
value_layout = (
cute.make_layout((1, copy_elems))
if major_mode == utils.LayoutEnum.ROW_MAJOR
else cute.make_layout((copy_elems, 1))
)
tiler_mn, layout_tv = cute.make_layout_tv(thread_layout, value_layout)
return cute.make_tiled_copy(atom_copy, layout_tv, tiler_mn)
def run_tensor_op_gemm(
a_major: str,
b_major: str,
c_major: str,
ab_dtype: Type[cutlass.Numeric],
c_dtype: Type[cutlass.Numeric],
acc_dtype: Type[cutlass.Numeric],
problem_shape: Tuple[int, int, int, int],
atom_layout_mnk: Tuple[int, int, int],
warmup_iterations: int = 2,
iterations: int = 100,
skip_ref_check: bool = False,
):
M, N, K, L = problem_shape
# Create and permute tensor A/B/C
def create_and_permute_tensor(l, mode0, mode1, is_mode0_major, dtype):
# is_mode0_major: (l, mode1, mode0) -> (mode0, mode1, l)
# else: (l, mode0, mode1) -> (mode0, mode1, l)
shape = (l, mode1, mode0) if is_mode0_major else (l, mode0, mode1)
permute_order = (2, 1, 0) if is_mode0_major else (1, 2, 0)
return (
torch.empty(*shape, dtype=torch.int32)
.random_(-2, 2)
.to(dtype=dtype)
.permute(permute_order)
.cuda()
)
a = create_and_permute_tensor(
L, M, K, a_major == "m", cutlass_torch.dtype(ab_dtype)
)
b = create_and_permute_tensor(
L, N, K, b_major == "n", cutlass_torch.dtype(ab_dtype)
)
c = create_and_permute_tensor(L, M, N, c_major == "m", cutlass_torch.dtype(c_dtype))
ref = torch.einsum("mkl,nkl->mnl", a, b).to(cutlass_torch.dtype(c_dtype))
tensor_op_gemm = TensorOpGemm(
ab_dtype,
c_dtype,
acc_dtype,
atom_layout_mnk,
)
# assume input is 16B aligned
a_tensor = (
from_dlpack(a, assumed_align=16)
.mark_layout_dynamic(leading_dim=(1 if a_major == "k" else 0))
.mark_compact_shape_dynamic(
mode=(1 if a_major == "k" else 0),
stride_order=(2, 0, 1) if a_major == "k" else (2, 1, 0),
divisibility=(128 // ab_dtype.width),
)
)
b_tensor = (
from_dlpack(b, assumed_align=16)
.mark_layout_dynamic(leading_dim=(1 if b_major == "k" else 0))
.mark_compact_shape_dynamic(
mode=(1 if b_major == "k" else 0),
stride_order=(2, 0, 1) if b_major == "k" else (2, 1, 0),
divisibility=(128 // ab_dtype.width),
)
)
c_tensor = (
from_dlpack(c, assumed_align=16)
.mark_layout_dynamic(leading_dim=(1 if c_major == "n" else 0))
.mark_compact_shape_dynamic(
mode=(1 if c_major == "n" else 0),
stride_order=(2, 0, 1) if c_major == "n" else (2, 1, 0),
divisibility=(128 // c_dtype.width),
)
)
print("Compiling kernel with cute.compile ...")
gemm = cute.compile(tensor_op_gemm, a_tensor, b_tensor, c_tensor)
print("Executing GEMM kernel...")
# Warmup
for _ in range(warmup_iterations):
gemm(a_tensor, b_tensor, c_tensor)
# Execute the kernel
for _ in range(iterations):
gemm(a_tensor, b_tensor, c_tensor)
if not skip_ref_check:
print("Verifying results...")
torch.testing.assert_close(c.cpu(), ref.cpu(), atol=1e-03, rtol=1e-05)
print("Results verified successfully!")
if __name__ == "__main__":
def parse_comma_separated_ints(s: str) -> Tuple[int, ...]:
try:
return tuple(int(x.strip()) for x in s.split(","))
except ValueError:
raise argparse.ArgumentTypeError(
"Invalid format. Expected comma-separated integers."
)
parser = argparse.ArgumentParser(
description="example of multistage block matmul with CuTe on GPU"
)
parser.add_argument(
"--mnkl", type=parse_comma_separated_ints, default=(112, 136, 40, 1)
)
parser.add_argument(
"--atom_layout_mnk", type=parse_comma_separated_ints, default=(2, 2, 1)
)
parser.add_argument(
"--ab_dtype",
type=cutlass.dtype,
choices=[cutlass.Float16],
default=cutlass.Float16,
)
parser.add_argument(
"--acc_dtype",
type=cutlass.dtype,
choices=[cutlass.Float32],
default=cutlass.Float32,
)
parser.add_argument(
"--c_dtype",
type=cutlass.dtype,
choices=[cutlass.Float16],
default=cutlass.Float16,
)
parser.add_argument("--a_major", choices=["k", "m"], default="m")
parser.add_argument("--b_major", choices=["k", "n"], default="n")
parser.add_argument("--c_major", choices=["n", "m"], default="n")
parser.add_argument("--warmup_iterations", default=2, type=int)
parser.add_argument("--iterations", default=100, type=int)
parser.add_argument("--skip_ref_check", action="store_true")
args = parser.parse_args()
print("Running Ampere tensor core GEMM example:")
run_tensor_op_gemm(
args.a_major,
args.b_major,
args.c_major,
args.ab_dtype,
args.c_dtype,
args.acc_dtype,
args.mnkl,
args.atom_layout_mnk,
args.warmup_iterations,
args.iterations,
args.skip_ref_check,
)
print("PASS")