Release v4.0.0 (#2294)
This commit is contained in:
@@ -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
@@ -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")
|
||||
Reference in New Issue
Block a user