CUTLASS 3.2.1 (#1113)
* Updates for 3.2.1 release. * Minor fix in gemm op profiler for raster order. * Add scheduler mapping for raster order in the kernels.
This commit is contained in:
@@ -31,6 +31,5 @@
|
||||
cutlass_example_add_executable(
|
||||
08_turing_tensorop_gemm
|
||||
turing_tensorop_gemm.cu
|
||||
DISABLE_TESTS ON
|
||||
)
|
||||
|
||||
|
||||
@@ -291,8 +291,8 @@ int run() {
|
||||
LayoutInputB,
|
||||
ElementOutput,
|
||||
LayoutOutput,
|
||||
ElementComputeEpilogue,
|
||||
ElementComputeEpilogue>
|
||||
int32_t,
|
||||
int32_t>
|
||||
gemm_device;
|
||||
|
||||
// Launch device reference gemm kernel
|
||||
@@ -355,4 +355,3 @@ int main() {
|
||||
|
||||
return run();
|
||||
}
|
||||
|
||||
|
||||
@@ -143,7 +143,6 @@ compare if the output from CUTLASS kernel is same as the reference implicit GEMM
|
||||
#include "cutlass/util/tensor_view_io.h"
|
||||
|
||||
#include "helper.h"
|
||||
|
||||
// The code section below describes datatype for input, output tensors and computation between
|
||||
// elements
|
||||
using ElementAccumulator = int32_t; // Data type of accumulator
|
||||
@@ -675,7 +674,6 @@ Result profile_convolution(Options const &options) {
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
int main(int argc, char const **args) {
|
||||
@@ -762,11 +760,7 @@ int main(int argc, char const **args) {
|
||||
Result::print_header(std::cout, options) << std::endl;
|
||||
result.print(std::cout, 1, options) << std::endl;
|
||||
}
|
||||
|
||||
return 0;
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -31,6 +31,5 @@
|
||||
cutlass_example_add_executable(
|
||||
12_gemm_bias_relu
|
||||
gemm_bias_relu.cu
|
||||
DISABLE_TESTS ON
|
||||
)
|
||||
|
||||
|
||||
@@ -220,7 +220,6 @@ bool run_fused_conv2d_fprop_optimized_s8_sm75_rf_res() {
|
||||
|
||||
return pass;
|
||||
}
|
||||
|
||||
int main() {
|
||||
|
||||
std::vector<bool (*)()>funcs = {
|
||||
@@ -229,10 +228,6 @@ int main() {
|
||||
};
|
||||
|
||||
return testRun(75, funcs, "conv int8 RF residency");
|
||||
|
||||
}
|
||||
|
||||
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
|
||||
@@ -39,7 +39,6 @@
|
||||
#include "device/b2b_implicit_gemm_convolution.h"
|
||||
#include "b2b_interleaved_conv2d_run.h"
|
||||
#include "test_run.h"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
cutlass::conv::Conv2dProblemSize conv2d_s8_sm75_problem_size_0 (
|
||||
@@ -219,20 +218,13 @@ bool run_fused_conv2d_fprop_optimized_s8_sm75_shmem() {
|
||||
|
||||
return pass;
|
||||
}
|
||||
|
||||
|
||||
int main() {
|
||||
|
||||
std::vector<bool (*)()>funcs = {
|
||||
&run_nonfused_conv2d_fprop_optimized_s8_sm75,
|
||||
&run_fused_conv2d_fprop_optimized_s8_sm75_shmem
|
||||
};
|
||||
|
||||
return testRun(75, funcs, "conv int8 shmem staging");
|
||||
|
||||
}
|
||||
|
||||
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
|
||||
@@ -195,7 +195,6 @@ bool run_fused_gemm_s8_rf_res() {
|
||||
return passed;
|
||||
|
||||
}
|
||||
|
||||
int main() {
|
||||
|
||||
std::vector<bool (*)()>funcs = {
|
||||
@@ -204,9 +203,6 @@ int main() {
|
||||
};
|
||||
|
||||
return testRun(75, funcs, "gemm int8 RF residency");
|
||||
|
||||
|
||||
}
|
||||
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -43,7 +43,6 @@
|
||||
#include "device/b2b_gemm.h"
|
||||
#include "b2b_interleaved_gemm_run.h"
|
||||
#include "test_run.h"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
cutlass::gemm::GemmCoord gemm_s8_sm75_problem_size_0(128*640, 64, 576);
|
||||
@@ -197,18 +196,13 @@ bool run_fused_gemm_s8_shmem() {
|
||||
return passed;
|
||||
|
||||
}
|
||||
|
||||
int main() {
|
||||
|
||||
std::vector<bool (*)()>funcs = {
|
||||
&run_nonfused_gemm_s8,
|
||||
&run_fused_gemm_s8_shmem
|
||||
};
|
||||
|
||||
return testRun(75, funcs, "gemm int8 shmem staing");
|
||||
|
||||
|
||||
}
|
||||
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -90,34 +90,6 @@ struct GroupedThreadblockSwizzle : detail::GroupedThreadblockSwizzleBase {
|
||||
}
|
||||
};
|
||||
|
||||
template <
|
||||
typename ThreadblockShape,
|
||||
typename LayoutC,
|
||||
cutlass::gemm::kernel::GroupScheduleMode GroupScheduleMode_ = cutlass::gemm::kernel::GroupScheduleMode::kDeviceOnly,
|
||||
int PrefetchTileCount = 128,
|
||||
int ThreadCount = PrefetchTileCount>
|
||||
struct GemmGroupedThreadblockSwizzle : GroupedThreadblockSwizzle<
|
||||
cutlass::gemm::kernel::GemmGroupedProblemVisitor<
|
||||
ThreadblockShape,
|
||||
GroupScheduleMode_,
|
||||
PrefetchTileCount,
|
||||
ThreadCount,
|
||||
platform::is_same<LayoutC, cutlass::layout::ColumnMajor>::value
|
||||
>
|
||||
> {
|
||||
using Base = GroupedThreadblockSwizzle<cutlass::gemm::kernel::GemmGroupedProblemVisitor<
|
||||
ThreadblockShape,
|
||||
GroupScheduleMode_,
|
||||
PrefetchTileCount,
|
||||
ThreadCount,
|
||||
platform::is_same<LayoutC, cutlass::layout::ColumnMajor>::value>>;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
GemmGroupedThreadblockSwizzle(typename Base::ProblemVisitor::Params& params,
|
||||
typename Base::ProblemVisitor::SharedStorage& shared_storage,
|
||||
int block_idx) : Base(params, shared_storage, block_idx) {}
|
||||
};
|
||||
|
||||
template <
|
||||
typename ThreadblockShape,
|
||||
typename LayoutC,
|
||||
|
||||
@@ -31,6 +31,7 @@
|
||||
|
||||
cutlass_example_add_executable(
|
||||
24_gemm_grouped
|
||||
gemm_grouped.cu
|
||||
gemm_grouped.cu
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -1,27 +1,4 @@
|
||||
# PyCUTLASS Examples
|
||||
|
||||
**NOTE:** This directory contains examples for PyCUTLASS, a Python library providing low-level
|
||||
building blocks for emitting CUTLASS C++ kernels. For examples using CUTLASS's Pythonic interface,
|
||||
see the [examples/python](/examples/python) directory.
|
||||
|
||||
Two types of examples are provided:
|
||||
* _Basic examples_: minimal examples that illustrate how to set up GEMMs, convolutions, and grouped GEMM operations
|
||||
* [_Customizable examples_](customizable): examples that allow one to specify a variety of template parameters for the given kernel
|
||||
|
||||
## Setting up the Python interface
|
||||
Please follow the instructions [here](/python/README.md#installation) to set up the PyCUTLASS.
|
||||
|
||||
## Running examples
|
||||
Each of the basic examples can be run as follows:
|
||||
```shell
|
||||
# Run the GEMM example
|
||||
python gemm.py
|
||||
|
||||
# Run the Conv2d example
|
||||
python conv2d.py
|
||||
|
||||
# Run the grouped GEMM example
|
||||
python gemm_grouped.py
|
||||
```
|
||||
|
||||
To run the customizable examples, refer to the README in the [customizable](customizable) directory.
|
||||
This directory contains deprecated examples for PyCUTLASS, a precursor to the CUTLASS Python interface.
|
||||
For examples of using CUTLASS's actively-maintained Pythonic interface, see the [examples/python](/examples/python) directory.
|
||||
|
||||
@@ -33,10 +33,14 @@
|
||||
Basic example of using the CUTLASS Python interface to run a 2d convolution
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import torch
|
||||
import numpy as np
|
||||
import sys
|
||||
print("This example is deprecated. Please see examples/python for examples of using "
|
||||
"the CUTLASS Python interface.")
|
||||
sys.exit(0)
|
||||
|
||||
import argparse
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
import cutlass_bindings
|
||||
import cutlass.backend as pycutlass
|
||||
|
||||
@@ -165,28 +165,3 @@ Example 7: GELU
|
||||
```python
|
||||
python gemm.py -i 16 8 16 -ta bfloat16 -tb bfloat16 -tc float32 -tacc float32 -m multiply_add -op TensorOp -b 64 128 64 -s 3 -w 2 2 1 -cc 80 -la ColumnMajor -aa 8 -lb ColumnMajor -ab 8 -lc RowMajor -ac 4 -te float32 -ep LinearCombination -sw IdentitySwizzle2 -p 512 256 128 -alpha 0.0 -beta 0.5 -gm GemmSplitKParallel -k 5 -bias -activ gelu
|
||||
```
|
||||
### Epilogue Visitor Tree
|
||||
Example 1:
|
||||
```python
|
||||
python gemm.py -i 16 8 8 -ta float32 -tb float32 -tc float32 -tacc float32 -m multiply_add_fast_bf16 -op TensorOp -b 128 128 32 -s 3 -w 2 2 1 -cc 80 -la RowMajor -aa 4 -lb ColumnMajor -ab 4 -lc RowMajor -ac 4 -te float32 -ep LinearCombination -epv RowBroadcast -sw IdentitySwizzle1 -p 512 256 128 -alpha 1.0 -beta 0.5 -gm Gemm -k 1
|
||||
```
|
||||
Example 2:
|
||||
```python
|
||||
python gemm.py -i 8 8 4 -ta float64 -tb float64 -tc float64 -tacc float64 -m multiply_add -op TensorOp -b 32 32 16 -s 4 -w 2 2 1 -cc 80 -la ColumnMajor -aa 1 -lb RowMajor -ab 1 -lc RowMajor -ac 1 -te float64 -ep LinearCombination -epv ColumnBroadcast -sw IdentitySwizzle1 -p 512 256 128 -alpha 1.0 -beta 0.5 -gm Gemm -k 1
|
||||
```
|
||||
Example 3:
|
||||
```python
|
||||
python gemm.py -i 16 8 16 -ta float16 -tb float16 -tc float32 -tacc float32 -m multiply_add -op TensorOp -b 128 128 32 -s 3 -w 2 2 1 -cc 80 -la ColumnMajor -aa 8 -lb RowMajor -ab 8 -lc RowMajor -ac 4 -te float32 -ep LinearCombination -epv RowReduction -sw IdentitySwizzle4 -p 512 256 128 -alpha 1.0 -beta 0.5 -gm Gemm -k 1
|
||||
```
|
||||
Example 4:
|
||||
```python
|
||||
python gemm.py -i 16 8 16 -ta bfloat16 -tb bfloat16 -tc float32 -tacc float32 -m multiply_add -op TensorOp -b 64 128 64 -s 3 -w 2 2 1 -cc 80 -la ColumnMajor -aa 8 -lb ColumnMajor -ab 8 -lc RowMajor -ac 4 -te float32 -ep LinearCombination -epv ColumnReduction -sw IdentitySwizzle2 -p 512 256 128 -alpha 1.0 -beta 0.5 -gm Gemm -k 1
|
||||
```
|
||||
Example 5:
|
||||
```python
|
||||
python gemm.py -i 16 8 8 -ta float32 -tb float32 -tc float32 -tacc float32 -m multiply_add_fast_bf16 -op TensorOp -b 128 128 32 -s 3 -w 2 2 1 -cc 80 -la RowMajor -aa 4 -lb ColumnMajor -ab 4 -lc RowMajor -ac 4 -te float32 -ep LinearCombination -epv RowReduction -sw BatchedIdentitySwizzle -p 512 256 128 -alpha 1.0 -beta 0.5 -gm Batched -k 1 -batch 3
|
||||
```
|
||||
Example 6:
|
||||
```python
|
||||
python gemm.py -i 16 8 8 -ta float32 -tb float32 -tc float32 -tacc float32 -m multiply_add_fast_bf16 -op TensorOp -b 128 128 32 -s 3 -w 2 2 1 -cc 80 -la RowMajor -aa 4 -lb ColumnMajor -ab 4 -lc RowMajor -ac 4 -te float32 -ep LinearCombination -epv ColumnBroadcast -sw BatchedIdentitySwizzle -p 512 256 128 -alpha 1.0 -beta 0.5 -gm Array -k 1 -batch 3
|
||||
```
|
||||
|
||||
@@ -29,13 +29,18 @@
|
||||
# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
#
|
||||
################################################################################
|
||||
|
||||
import sys
|
||||
print("This example is deprecated. Please see examples/python for examples of using "
|
||||
"the CUTLASS Python interface.")
|
||||
sys.exit(0)
|
||||
|
||||
import numpy as np
|
||||
import cutlass.backend as pycutlass
|
||||
from cutlass.backend import *
|
||||
from cutlass.backend.utils.device import device_cc
|
||||
from cutlass.backend.conv2d_operation import *
|
||||
from cutlass.backend.utils.reference_model import Conv2dReferenceModule
|
||||
import sys
|
||||
import torch.nn.functional as F
|
||||
|
||||
import argparse
|
||||
|
||||
@@ -29,13 +29,18 @@
|
||||
# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
#
|
||||
################################################################################
|
||||
|
||||
import sys
|
||||
print("This example is deprecated. Please see examples/python for examples of using "
|
||||
"the CUTLASS Python interface.")
|
||||
sys.exit(0)
|
||||
|
||||
import numpy as np
|
||||
import cutlass.backend as pycutlass
|
||||
from cutlass.backend import *
|
||||
from cutlass.backend.utils.device import device_cc
|
||||
import cutlass_bindings
|
||||
from bfloat16 import bfloat16
|
||||
import sys
|
||||
|
||||
import argparse
|
||||
|
||||
@@ -100,8 +105,6 @@ parser.add_argument("-te", "--element_epilogue", default="float32", type=str,
|
||||
parser.add_argument("-ep", "--epilogue_functor", default="LinearCombination",
|
||||
type=str, choices=['LinearCombination', 'FastLinearCombinationClamp', 'LinearCombinationClamp'],
|
||||
help="This option describes the epilogue part of the kernel")
|
||||
parser.add_argument("-epv", "--epilogue_visitor", default=None,
|
||||
type=str, choices=['RowReduction', 'ColumnReduction', 'RowBroadcast', 'ColumnBroadcast'], help="epilogue visitor for more complex epilogues")
|
||||
# swizzling
|
||||
parser.add_argument("-sw", "--swizzling_functor", default="IdentitySwizzle1", type=str, choices=[
|
||||
"IdentitySwizzle1", "IdentitySwizzle2", "IdentitySwizzle4", "IdentitySwizzle8", "HorizontalSwizzle", "BatchedIdentitySwizzle"],
|
||||
@@ -193,71 +196,10 @@ else:
|
||||
|
||||
swizzling_functor = getattr(cutlass_bindings, args.swizzling_functor)
|
||||
|
||||
visitor = args.epilogue_visitor is not None
|
||||
|
||||
if args.epilogue_visitor == "ColumnReduction":
|
||||
class ColumnReduction_(EpilogueVisitTree):
|
||||
def __call__(
|
||||
self, accum: 'tensor', c: 'tensor',
|
||||
alpha: 'scalar', beta: 'scalar'):
|
||||
#
|
||||
D = alpha * accum + beta * c
|
||||
reduction = reduction_op(D, "column", "Add", args.threadblock_shape[0])
|
||||
return D, reduction
|
||||
epilogue_functor = ColumnReduction_(
|
||||
epilogue_functor, tile_description, math_inst.element_accumulator,
|
||||
C.alignment, element_epilogue, C.element)
|
||||
epilogue_functor.initialize()
|
||||
elif args.epilogue_visitor == "RowReduction":
|
||||
class RowReduction_(EpilogueVisitTree):
|
||||
def __call__(
|
||||
self, accum: 'tensor', c: 'tensor',
|
||||
alpha: 'scalar', beta: 'scalar'):
|
||||
#
|
||||
D = alpha * accum + tanh.numpy(beta * c)
|
||||
reduction = reduction_op(D, "row", "Add", args.threadblock_shape[1])
|
||||
return D, reduction
|
||||
epilogue_functor = RowReduction_(
|
||||
epilogue_functor, tile_description, math_inst.element_accumulator,
|
||||
C.alignment, element_epilogue, C.element)
|
||||
epilogue_functor.initialize()
|
||||
|
||||
elif args.epilogue_visitor == "RowBroadcast":
|
||||
class RowBroadcast_(EpilogueVisitTree):
|
||||
def __call__(
|
||||
self, accum: 'tensor', c: 'tensor',
|
||||
vector: 'row', alpha: 'scalar', beta: 'scalar'):
|
||||
#
|
||||
T = accum + vector
|
||||
scale_T = alpha * T
|
||||
Z = relu.numpy(scale_T + beta * c)
|
||||
return Z, T
|
||||
epilogue_functor = RowBroadcast_(
|
||||
epilogue_functor, tile_description, math_inst.element_accumulator,
|
||||
C.alignment, element_epilogue, C.element)
|
||||
epilogue_functor.initialize()
|
||||
elif args.epilogue_visitor == "ColumnBroadcast":
|
||||
class ColumnBroadcast_(EpilogueVisitTree):
|
||||
def __call__(
|
||||
self, accum: 'tensor', c: 'tensor',
|
||||
vector: 'column', alpha: 'scalar', beta: 'scalar'):
|
||||
#
|
||||
T = accum + vector
|
||||
scale_T = leaky_relu.numpy(alpha * T, 0.2)
|
||||
Z = scale_T + beta * c
|
||||
return Z, T
|
||||
epilogue_functor = ColumnBroadcast_(
|
||||
epilogue_functor, tile_description, math_inst.element_accumulator,
|
||||
C.alignment, element_epilogue, C.element)
|
||||
epilogue_functor.initialize()
|
||||
else:
|
||||
epilogue_functor = epilogue_functor
|
||||
|
||||
operation = GemmOperationUniversal(
|
||||
arch=args.compute_capability, tile_description=tile_description,
|
||||
A=A, B=B, C=C,
|
||||
epilogue_functor=epilogue_functor, swizzling_functor=swizzling_functor,
|
||||
visitor=visitor
|
||||
epilogue_functor=epilogue_functor, swizzling_functor=swizzling_functor
|
||||
)
|
||||
|
||||
if args.print_cuda:
|
||||
@@ -347,38 +289,7 @@ tensor_D = np.zeros(
|
||||
shape=(args.batch * problem_size.m() * problem_size.n(),)
|
||||
).astype(getattr(np, args.element_c))
|
||||
|
||||
if args.epilogue_visitor == "RowReduction":
|
||||
cta_n = args.threadblock_shape[1]
|
||||
num_cta_n = (problem_size.n() + cta_n - 1) // cta_n
|
||||
reduction = np.zeros(shape=(args.batch * problem_size.m() * num_cta_n,), dtype=getattr(np, args.element_c))
|
||||
output_op = operation.epilogue_type(
|
||||
D=tensor_D, alpha=args.alpha, beta=args.beta, c=tensor_C, reduction=reduction, problem_size=[problem_size.m(), problem_size.n()]
|
||||
)
|
||||
elif args.epilogue_visitor == "ColumnReduction":
|
||||
cta_m = args.threadblock_shape[0]
|
||||
num_cta_m = (problem_size.m() + cta_m - 1) // cta_m
|
||||
reduction = np.zeros(shape=(args.batch * problem_size.n() * num_cta_m,), dtype=getattr(np, args.element_c))
|
||||
output_op = operation.epilogue_type(
|
||||
D=tensor_D, alpha=args.alpha, beta=args.beta, c=tensor_C, reduction=reduction, problem_size=[problem_size.m(), problem_size.n()]
|
||||
)
|
||||
elif args.epilogue_visitor == "RowBroadcast":
|
||||
vector = np.ceil(
|
||||
np.random.uniform(low=-8.5, high=7.5, size=(args.batch, 1, problem_size.n()))
|
||||
).astype(getattr(np, args.element_c))
|
||||
tensor_t = np.empty_like(tensor_D)
|
||||
output_op = operation.epilogue_type(
|
||||
c=tensor_C, vector=vector, alpha=args.alpha, beta=args.beta, Z=tensor_D, T=tensor_t, problem_size=[problem_size.m(), problem_size.n()]
|
||||
)
|
||||
elif args.epilogue_visitor == "ColumnBroadcast":
|
||||
vector = np.ceil(
|
||||
np.random.uniform(low=-8.5, high=7.5, size=(args.batch, problem_size.m(), 1))
|
||||
).astype(getattr(np, args.element_c))
|
||||
tensor_t = np.empty_like(tensor_D)
|
||||
output_op = operation.epilogue_type(
|
||||
c=tensor_C, vector=vector, alpha=args.alpha, beta=args.beta, Z=tensor_D, T=tensor_t, problem_size=[problem_size.m(), problem_size.n()]
|
||||
)
|
||||
else:
|
||||
output_op = operation.epilogue_type(*([args.alpha, args.beta] + args.activation_args))
|
||||
output_op = operation.epilogue_type(*([args.alpha, args.beta] + args.activation_args))
|
||||
|
||||
arguments = GemmArguments(
|
||||
operation=operation, problem_size=problem_size,
|
||||
@@ -411,38 +322,8 @@ reference = ReferenceModule(A, B, C)
|
||||
tensor_D_ref = reference.run(
|
||||
tensor_A, tensor_B, tensor_C, problem_size, args.alpha, args.beta, args.bias, args.batch)
|
||||
|
||||
if args.epilogue_visitor in ["RowBroadcast", "ColumnBroadcast"]:
|
||||
tensor_D_ref = (tensor_D_ref.reshape((args.batch, problem_size.m(), problem_size.n())) + vector).flatten()
|
||||
tensor_D_ref = getattr(pycutlass, args.activation_function).numpy(*([tensor_D_ref,] + args.activation_args))
|
||||
|
||||
if args.epilogue_visitor in ["RowReduction", "ColumnReduction"]:
|
||||
output_op.sync()
|
||||
accum_ref = reference.run(
|
||||
tensor_A, tensor_B, tensor_C, problem_size, 1.0, 0.0, args.bias, args.batch)
|
||||
tensor_D_ref, reduction_ref = epilogue_functor(
|
||||
accum_ref.reshape((args.batch, problem_size.m(), problem_size.n())),
|
||||
tensor_C.reshape((args.batch, problem_size.m(), problem_size.n())),
|
||||
args.alpha, args.beta
|
||||
)
|
||||
tensor_D_ref = tensor_D_ref.flatten()
|
||||
reduction_ref = reduction_ref.flatten()
|
||||
assert np.allclose(reduction_ref, reduction, atol=1e-2)
|
||||
|
||||
elif args.epilogue_visitor in ["RowBroadcast", "ColumnBroadcast"]:
|
||||
output_op.sync()
|
||||
accum_ref = reference.run(
|
||||
tensor_A, tensor_B, tensor_C, problem_size, 1.0, 0.0, args.bias, args.batch)
|
||||
|
||||
tensor_D_ref, tensor_T_ref = epilogue_functor(
|
||||
accum_ref.reshape((args.batch, problem_size.m(), problem_size.n())),
|
||||
tensor_C.reshape((args.batch, problem_size.m(), problem_size.n())),
|
||||
vector, args.alpha, args.beta)
|
||||
|
||||
tensor_D_ref = tensor_D_ref.flatten()
|
||||
tensor_T_ref = tensor_T_ref.flatten()
|
||||
|
||||
assert np.array_equal(tensor_t, tensor_T_ref)
|
||||
|
||||
try:
|
||||
assert np.array_equal(tensor_D, tensor_D_ref)
|
||||
except:
|
||||
|
||||
@@ -29,12 +29,17 @@
|
||||
# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
#
|
||||
################################################################################
|
||||
|
||||
import sys
|
||||
print("This example is deprecated. Please see examples/python for examples of using "
|
||||
"the CUTLASS Python interface.")
|
||||
sys.exit(0)
|
||||
|
||||
import numpy as np
|
||||
import cutlass.backend as pycutlass
|
||||
from cutlass.backend import *
|
||||
from cutlass.backend.utils.device import device_cc
|
||||
import csv
|
||||
import sys
|
||||
|
||||
import argparse
|
||||
|
||||
|
||||
@@ -33,9 +33,13 @@
|
||||
Basic example of using the CUTLASS Python interface to run a GEMM
|
||||
"""
|
||||
|
||||
import sys
|
||||
print("This example is deprecated. Please see examples/python for examples of using "
|
||||
"the CUTLASS Python interface.")
|
||||
sys.exit(0)
|
||||
|
||||
import argparse
|
||||
import numpy as np
|
||||
import sys
|
||||
|
||||
import cutlass_bindings
|
||||
import cutlass.backend as pycutlass
|
||||
|
||||
@@ -33,9 +33,13 @@
|
||||
Basic example of using the CUTLASS Python interface to run a grouped GEMM
|
||||
"""
|
||||
|
||||
import sys
|
||||
print("This example is deprecated. Please see examples/python for examples of using "
|
||||
"the CUTLASS Python interface.")
|
||||
sys.exit(0)
|
||||
|
||||
import argparse
|
||||
import numpy as np
|
||||
import sys
|
||||
|
||||
import cutlass_bindings
|
||||
import cutlass.backend as pycutlass
|
||||
|
||||
@@ -434,14 +434,6 @@ class gen_device:
|
||||
" if (result != cudaSuccess) {\n" + \
|
||||
" return Status::kErrorInternal;\n" + \
|
||||
" }\n" + \
|
||||
"\n" + \
|
||||
" result = cudaFuncSetAttribute(\n" + \
|
||||
" Kernel<B2bGemmKernel>,\n" + \
|
||||
" cudaFuncAttributePreferredSharedMemoryCarveout, 100);\n" + \
|
||||
"\n" + \
|
||||
" if (result != cudaSuccess) {\n" + \
|
||||
" return Status::kErrorInternal;\n" + \
|
||||
" }\n" + \
|
||||
" }\n" + \
|
||||
" cutlass::Kernel<B2bGemmKernel><<<grid, block, smem_size, stream>>>(params_);\n" + \
|
||||
" result = cudaGetLastError();\n" + \
|
||||
|
||||
@@ -83,6 +83,10 @@
|
||||
#include "cutlass/util/reference/host/tensor_fill.h"
|
||||
#include "cutlass/util/tensor_view_io.h"
|
||||
|
||||
#include "cutlass/epilogue/threadblock/fusion/visitors.hpp"
|
||||
#include "cutlass/gemm/kernel/default_gemm_universal_with_visitor.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
|
||||
#include "helper.h"
|
||||
|
||||
|
||||
@@ -120,6 +124,7 @@ using ThreadblockShape = cutlass::gemm::GemmShape<128, 128, 32>; // Threadb
|
||||
using WarpShape = cutlass::gemm::GemmShape<64, 64, 32>; // Warp-level tile size (concept: GemmShape)
|
||||
using InstructionShape = cutlass::gemm::GemmShape<16, 8, 16>; // Instruction-level tile size (concept: GemmShape)
|
||||
constexpr int NumStages = 4; // Number of global->shared pipeline stages used in the GEMM mainloop
|
||||
constexpr int EVTEpilogueStages = 1; // Number of epilogue stages in EVT
|
||||
|
||||
// Residual block configuration
|
||||
|
||||
@@ -166,23 +171,93 @@ using DeviceGemmBasic = cutlass::gemm::device::GemmUniversalWithBroadcast<
|
||||
AlignmentA,
|
||||
AlignmentB>;
|
||||
|
||||
// StreamK device GEMM implementation type
|
||||
using DeviceGemmStreamK = cutlass::gemm::device::GemmUniversalStreamkWithBroadcast<
|
||||
ElementA, LayoutA,
|
||||
ElementB, LayoutB,
|
||||
ElementC, LayoutC,
|
||||
// StreamK device GEMM implementation type with EVT
|
||||
using namespace cute;
|
||||
|
||||
using OutputTileThreadMap = cutlass::epilogue::threadblock::OutputTileThreadLayout<
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
ElementC,
|
||||
AlignmentC,
|
||||
EVTEpilogueStages
|
||||
>;
|
||||
|
||||
using Accum = cutlass::epilogue::threadblock::VisitorAccFetch;
|
||||
|
||||
using Bias = cutlass::epilogue::threadblock::VisitorRowBroadcast<
|
||||
OutputTileThreadMap, ElementC,
|
||||
cute::Stride<_0, _1, int32_t> // StrideMNL
|
||||
>;
|
||||
|
||||
using C1 = cutlass::epilogue::threadblock::VisitorAuxLoad<
|
||||
OutputTileThreadMap, ElementC,
|
||||
cute::Stride<int64_t, _1, int64_t> // StrideMNL
|
||||
>;
|
||||
|
||||
using C2 = cutlass::epilogue::threadblock::VisitorAuxLoad<
|
||||
OutputTileThreadMap, ElementC,
|
||||
cute::Stride<int64_t, _1, int64_t> // StrideMNL
|
||||
>;
|
||||
|
||||
using Compute0 = cutlass::epilogue::threadblock::VisitorCompute<
|
||||
cutlass::plus, ElementCompute, ElementCompute,
|
||||
cutlass::FloatRoundStyle::round_to_nearest
|
||||
>;
|
||||
|
||||
using EVTCompute0 = cutlass::epilogue::threadblock::Sm80EVT<
|
||||
Compute0,
|
||||
Accum,
|
||||
Bias>;
|
||||
|
||||
using Compute1 = cutlass::epilogue::threadblock::VisitorCompute<
|
||||
cutlass::plus, ElementCompute, ElementCompute,
|
||||
cutlass::FloatRoundStyle::round_to_nearest
|
||||
>;
|
||||
|
||||
using EVTCompute1 = cutlass::epilogue::threadblock::Sm80EVT<
|
||||
Compute1,
|
||||
EVTCompute0,
|
||||
C1>;
|
||||
|
||||
using Compute2 = cutlass::epilogue::threadblock::VisitorCompute<
|
||||
cutlass::plus, ElementOutput, ElementCompute,
|
||||
cutlass::FloatRoundStyle::round_to_nearest
|
||||
>;
|
||||
|
||||
using EVTCompute2 = cutlass::epilogue::threadblock::Sm80EVT<
|
||||
Compute2,
|
||||
EVTCompute1,
|
||||
C2>;
|
||||
|
||||
using D = cutlass::epilogue::threadblock::VisitorAuxStore<
|
||||
OutputTileThreadMap, ElementOutput, cutlass::FloatRoundStyle::round_to_nearest,
|
||||
cute::Stride<int64_t, _1, int64_t> // StrideMNL
|
||||
>;
|
||||
|
||||
using EVTD = cutlass::epilogue::threadblock::Sm80EVT<
|
||||
D,
|
||||
EVTCompute2>;
|
||||
|
||||
using EVTKernelStreamK =
|
||||
typename cutlass::gemm::kernel::DefaultGemmWithVisitor<
|
||||
ElementA, LayoutA, cutlass::ComplexTransform::kNone, AlignmentA,
|
||||
ElementB, LayoutB, cutlass::ComplexTransform::kNone, AlignmentB,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementAccumulator,
|
||||
OperatorClass,
|
||||
ArchTag,
|
||||
ElementCompute,
|
||||
cutlass::arch::OpClassTensorOp,
|
||||
cutlass::arch::Sm80,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOp,
|
||||
EVTD,
|
||||
cutlass::gemm::threadblock::ThreadblockSwizzleStreamK,
|
||||
NumStages,
|
||||
AlignmentA,
|
||||
AlignmentB>;
|
||||
cutlass::arch::OpMultiplyAdd,
|
||||
EVTEpilogueStages
|
||||
>::GemmKernel;
|
||||
|
||||
using DeviceGemmStreamK = cutlass::gemm::device::GemmUniversalAdapter<EVTKernelStreamK>;
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Testbed utility types
|
||||
@@ -360,36 +435,41 @@ typename DeviceGemmStreamK::Arguments args_from_options(
|
||||
cutlass::HostTensor<ElementC, LayoutC> &tensor_Vector/*,
|
||||
cutlass::HostTensor<ElementC, LayoutC> &tensor_Tensor*/
|
||||
)
|
||||
{
|
||||
{
|
||||
typename EVTD::Arguments callback_args{
|
||||
{
|
||||
{
|
||||
{
|
||||
{}, // Accum
|
||||
{tensor_Vector.device_data(), ElementC(0), {_0{}, _1{}, int32_t(options.problem_size.n())}}, // Bias
|
||||
{} // Compute0
|
||||
}, // EVTCompute0
|
||||
{tensor_c1.device_data(), ElementC(0), {options.problem_size.n(), _1{}, options.problem_size.mn().product()}}, // C1
|
||||
{} // Compute1
|
||||
}, // EVTCompute1
|
||||
{tensor_c2.device_data(), ElementC(0), {options.problem_size.n(), _1{}, options.problem_size.mn().product()}}, // C2
|
||||
{} // Compute2
|
||||
}, // EVTCompute2
|
||||
{tensor_d.device_data(), {options.problem_size.n(), _1{}, options.problem_size.mn().product()}}, // D
|
||||
}; // EVTD
|
||||
|
||||
return typename DeviceGemmStreamK::Arguments(
|
||||
cutlass::gemm::GemmUniversalMode::kGemm, // universal mode
|
||||
options.problem_size, // problem_size
|
||||
options.split_k_factor, // batch count / splitk slices
|
||||
{ // epilogue parameters
|
||||
ElementAccumulator(options.alpha),
|
||||
ElementAccumulator(options.beta)
|
||||
},
|
||||
callback_args, // argument of EVT callbacks
|
||||
tensor_a.device_data(), // ptr_A
|
||||
tensor_b.device_data(), // ptr_B
|
||||
tensor_c1.device_data(), // ptr_C1
|
||||
tensor_c2.device_data(), // ptr_C2
|
||||
tensor_d.device_data(), // ptr_D
|
||||
tensor_Vector.device_data(), // ptr_Vector
|
||||
/* tensor_Tensor.device_data(), */nullptr,// ptr_Tensor // We're not storing Tensor
|
||||
nullptr, // ptr_C (unused)
|
||||
nullptr, // ptr_D (unused)
|
||||
options.problem_size.mk().product(), // batch_stride_A
|
||||
options.problem_size.nk().product(), // batch_stride_B
|
||||
options.problem_size.mn().product(), // batch_stride_C1
|
||||
options.problem_size.mn().product(), // batch_stride_C2
|
||||
options.problem_size.mn().product(), // batch_stride_D
|
||||
options.problem_size.mn().product(), // batch_stride_Vector
|
||||
options.problem_size.mn().product(), // batch_stride_Tensor
|
||||
0, // batch_stride_C (unused)
|
||||
0, // batch_stride_D (unused)
|
||||
tensor_a.layout().stride(0), // stride_a
|
||||
tensor_b.layout().stride(0), // stride_b
|
||||
tensor_c1.layout().stride(0), // stride_c1
|
||||
tensor_c2.layout().stride(0), // stride_c2
|
||||
tensor_d.layout().stride(0), // stride_d
|
||||
/*tensor_Vector.layout().stride(0)*/0, // stride_Vector // Vector stride is always 0
|
||||
/*tensor_Tensor.layout().stride(0)*/0, // stride_Tensor // We're not storing Tensor
|
||||
0, // stride_c (unused)
|
||||
0, // stride_d (unused)
|
||||
options.avail_sms); // avail_sms
|
||||
}
|
||||
|
||||
|
||||
@@ -526,7 +526,8 @@ struct ExampleRunner
|
||||
|
||||
// Forward calls via lambda to avoid specifying template arguments
|
||||
auto gather_call = [](auto&&... args){ gather(static_cast<decltype(args)&&>(args)...); };
|
||||
auto scatter_call = [](auto&&... args){ scatter(static_cast<decltype(args)&&>(args)...); };
|
||||
// MSVC doesn't count use inside a false "if constexpr" branch.
|
||||
[[maybe_unused]] auto scatter_call = [](auto&&... args){ scatter(static_cast<decltype(args)&&>(args)...); };
|
||||
|
||||
if constexpr (DoGatherA) {
|
||||
run_gather(gather_call, tensor_a, tensor_a_gathered, arguments.gather_A, problem_size.batch(), stride_A);
|
||||
|
||||
@@ -58,7 +58,7 @@ public:
|
||||
// Type Aliases
|
||||
//
|
||||
using ProblemShape = ProblemShape_;
|
||||
using TileScheduleTag = TileScheduler_;
|
||||
using TileSchedulerTag = TileScheduler_;
|
||||
using TileScheduler = TileScheduler_;
|
||||
static_assert(rank(ProblemShape{}) == 3 or rank(ProblemShape{}) == 4,
|
||||
"ProblemShape{} should be <M,N,K> or <M,N,K,L>");
|
||||
|
||||
@@ -161,7 +161,7 @@ using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
using EpilogueOutputOp = typename Gemm::EpilogueOutputOp;
|
||||
using ElementScalar = typename EpilogueOutputOp::ElementScalar;
|
||||
using ElementAmax = typename EpilogueOutputOp::ElementAmax;
|
||||
using ActivationFunctor = typename EpilogueOutputOp::ActivationFn<ElementCompute>;
|
||||
using ActivationFunctor = typename EpilogueOutputOp::ActivationFn;
|
||||
|
||||
using StrideA = typename Gemm::GemmKernel::StrideA;
|
||||
using StrideB = typename Gemm::GemmKernel::StrideB;
|
||||
|
||||
@@ -7,9 +7,7 @@
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"# Basic example of using the CUTLASS Python interface\n",
|
||||
"This notebook walks through a basic example of using the CUTLASS Python interface to declare, compile, and run GEMMs.\n",
|
||||
"\n",
|
||||
"[](https://colab.research.google.com/github/NVIDIA/cutlass/tree/master/examples/00_basic_gemm.ipynb)\n"
|
||||
"This notebook walks through a basic example of using the CUTLASS Python interface to declare, compile, and run GEMMs.\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
|
||||
@@ -7,9 +7,7 @@
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"# Example of using elementwise activation functions in the CUTLASS Python interface\n",
|
||||
"This notebook walks through a basic example of using the CUTLASS Python interface to declare, compile, and run GEMMs with different epilogues.\n",
|
||||
"\n",
|
||||
"[](https://colab.research.google.com/github/NVIDIA/cutlass/tree/master/examples/00_basic_gemm.ipynb)"
|
||||
"This notebook walks through a basic example of using the CUTLASS Python interface to declare, compile, and run GEMMs with different epilogues.\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
|
||||
@@ -10,8 +10,6 @@
|
||||
"This notebook walks through a basic example of using the CUTLASS Python interface to declare\n",
|
||||
"a grouped GEMM kernel and export it as a PyTorch CUDA extension. Note that GEMM and Conv2d can also be exported as PyTorch CUDA extensions. \n",
|
||||
"\n",
|
||||
"[](https://colab.research.google.com/github/NVIDIA/cutlass/tree/master/examples/00_basic_gemm.ipynb)\n",
|
||||
"\n",
|
||||
"## Background on grouped GEMM\n",
|
||||
"Grouped GEMM enables one to execute a set of GEMMs (each with potentially different sizes and strides)\n",
|
||||
"in a single CUDA kernel. It can be thought of as a generalized version of a pointer-array GEMM,\n",
|
||||
|
||||
221
examples/python/04_epilogue_visitor.ipynb
Normal file
221
examples/python/04_epilogue_visitor.ipynb
Normal file
@@ -0,0 +1,221 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"attachments": {},
|
||||
"cell_type": "markdown",
|
||||
"id": "5d24a692",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"# Example of using epilogue visitor in the CUTLASS Python interface\n",
|
||||
"This notebook walks through a basic example of using the CUTLASS Python interface to declare, compile, and run GEMMs with different epilogues through CUTLASS Epilogue Visitor."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "3ca993fe",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"We first import various packages needed for the example, construct the input and output tensors that will be used in our example."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "63a70a3c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import torch\n",
|
||||
"import cutlass\n",
|
||||
"from cutlass.epilogue import relu\n",
|
||||
"from cutlass import Tensor as FakeTensor\n",
|
||||
"from cutlass.profiler import CUDAEventProfiler\n",
|
||||
"\n",
|
||||
"# This controls whether ther C++ GEMM declaration will be printed at each step. Set to `false` to\n",
|
||||
"# omit this information.\n",
|
||||
"print_module = True\n",
|
||||
"\n",
|
||||
"# The Epilogue Visitor feature currently only works for SM80 and 90\n",
|
||||
"from cutlass.backend.utils.device import device_cc\n",
|
||||
"if device_cc() not in [80, 90]:\n",
|
||||
" import sys\n",
|
||||
" sys.exit()\n",
|
||||
"\n",
|
||||
"m = 16384\n",
|
||||
"n = m\n",
|
||||
"k = 512\n",
|
||||
"\n",
|
||||
"type_A = torch.float16\n",
|
||||
"type_B = torch.float16\n",
|
||||
"type_C = torch.float16\n",
|
||||
"type_D = torch.float16\n",
|
||||
"\n",
|
||||
"torch.manual_seed(2023)\n",
|
||||
"scope_min = -4\n",
|
||||
"scope_max = 4\n",
|
||||
"tensor_A = torch.ceil(torch.empty(size=(m, k), dtype=type_A, device=\"cuda\").uniform_(scope_min, scope_max))\n",
|
||||
"tensor_B = torch.ceil(torch.empty(size=(k, n), dtype=type_B, device=\"cuda\").uniform_(scope_min, scope_max))\n",
|
||||
"tensor_C = torch.ceil(torch.empty(size=(m, n), dtype=type_C, device=\"cuda\").uniform_(scope_min, scope_max))\n",
|
||||
"tensor_D = torch.zeros_like(tensor_C)\n",
|
||||
"\n",
|
||||
"plan = cutlass.op.Gemm(element=torch.float16, layout=cutlass.LayoutType.RowMajor, element_accumulator=torch.float32)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "1eb0d95b",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Define the epilogue visitor functor\n",
|
||||
"The epilogue functor can be defined as a simple Python function and a set of example tensors for inputs and outputs. The example below illustrates a complex epilogue under the directed acyclic graph structure (`F` is used twice). The epilogue takes source tensors in different ranks: `alpha`, `beta` are scalars, `bias` is a column vector to broadcast, and `C`, `aux` are matrices. It contains various math operations from basic arithmatic operations and built-in callable functions like `relu`. It also accomodates multiple outputs `D` and `F`. Note that there are some restrictions on syntax.\n",
|
||||
"* Each named variable must be assigned exactly once and defined before it it used.\n",
|
||||
"* Reserved names: `accum`, `C`, and `D` are reserved for accumulator, tensor_C, and tensor_D.\n",
|
||||
"* Return values must be a named variable.\n",
|
||||
"\n",
|
||||
"The example tensors is a dictionary with tensor names as keys and reference tensors as values. The reference tensors can be `float`, `torch.Tensor`, `numpy.ndarray`, or our `FakeTensor`. They provides the shape and data type information of the inputs and outputs of the epilogue.\n",
|
||||
"\n",
|
||||
"The epilogue can be generated simply through `cutlass.evt.trace(<epilogue function>, <example_tensors>)`."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "8d257833",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Define epilogue visitor\n",
|
||||
"def example_epilogue(accum, alpha, C, beta, aux, bias):\n",
|
||||
" F = alpha * accum + (beta * C + aux)\n",
|
||||
" E = relu(F + 1) + bias\n",
|
||||
" D = E + F\n",
|
||||
" return D, F\n",
|
||||
"\n",
|
||||
"# Construct inputs and outputs\n",
|
||||
"alpha = 0.5\n",
|
||||
"beta = 0.5\n",
|
||||
"aux = torch.ceil(torch.empty(size=(m, n), dtype=type_C, device=\"cuda\").uniform_(scope_min, scope_max))\n",
|
||||
"bias = torch.ceil(torch.empty(size=(m, 1), dtype=type_C, device=\"cuda\").uniform_(scope_min, scope_max))\n",
|
||||
"tensor_F = torch.zeros_like(tensor_D)\n",
|
||||
"examples_tensors = {\n",
|
||||
" \"accum\": FakeTensor(element=torch.float32, shape=(m, n), layout_tag=cutlass.LayoutType.RowMajor),\n",
|
||||
" \"alpha\": alpha,\n",
|
||||
" \"C\": tensor_C,\n",
|
||||
" \"beta\": beta,\n",
|
||||
" \"aux\": aux,\n",
|
||||
" \"bias\": bias,\n",
|
||||
" \"D\": tensor_D,\n",
|
||||
" \"F\": tensor_F\n",
|
||||
"}\n",
|
||||
"\n",
|
||||
"# Trace the epilogue visitor\n",
|
||||
"epilogue_visitor = cutlass.epilogue.trace(example_epilogue, examples_tensors)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "54961694",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Run a GEMM with the epilogue visitor functor\n",
|
||||
"The `epilogue_visitor` can be used by setting the plan's `epilogue_visitor` field. The arguments for the epilogue visitor are provided as a `dict` through the `visitor_args` keyword argument."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "5fe49443",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"visitor_args = {\n",
|
||||
" \"alpha\": alpha, \"C\": tensor_C, \"beta\": beta, \n",
|
||||
" \"aux\": aux, \"bias\": bias, \"D\": tensor_D, \"F\": tensor_F\n",
|
||||
"}\n",
|
||||
"\n",
|
||||
"plan.epilogue_visitor = epilogue_visitor\n",
|
||||
"plan.run(\n",
|
||||
" tensor_A, tensor_B, tensor_C, tensor_D, \n",
|
||||
" visitor_args=visitor_args, print_module=print_module)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "455d0a37",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"The epilogue function `example_epilogue` can be used as a reference function. We can now verify the results simply with"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "e32e7798",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"class TorchReference(torch.nn.Module):\n",
|
||||
" def forward(self, A, B, alpha, C, beta, aux, bias):\n",
|
||||
" accum = torch.matmul(A, B)\n",
|
||||
" return example_epilogue(accum, alpha, C, beta, aux, bias)\n",
|
||||
"\n",
|
||||
"torch_reference = TorchReference()\n",
|
||||
"if hasattr(torch, \"compile\"):\n",
|
||||
" # If the torch.compile feature is available\n",
|
||||
" torch_reference = torch.compile(torch_reference)\n",
|
||||
"\n",
|
||||
"tensor_D_ref, tensor_F_ref = torch_reference(tensor_A, tensor_B, alpha, tensor_C, beta, aux, bias)\n",
|
||||
"\n",
|
||||
"assert torch.equal(tensor_D, tensor_D_ref)\n",
|
||||
"assert torch.equal(tensor_F, tensor_F_ref)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "b69e441f",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"The performance of CUTLASS fused kernel can be profiled with"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "8db92150",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"warmup_iterations = 10\n",
|
||||
"profile_iterations = 50\n",
|
||||
"# Profile CUTLASS fused kernel\n",
|
||||
"duration = CUDAEventProfiler(\n",
|
||||
" plan, warmup_iterations, profile_iterations,\n",
|
||||
" tensor_A, tensor_B, tensor_C, tensor_D, \n",
|
||||
" visitor_args=visitor_args)()\n",
|
||||
"\n",
|
||||
"print(f\"CUTLASS duration: {duration:.2f} ms\")"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.10"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -16,3 +16,7 @@
|
||||
* [03_basic_conv2d](/examples/python/03_basic_conv2d.ipynb)
|
||||
|
||||
Shows how to declare, configure, compile, and run a CUTLASS Conv2d using the Python interface
|
||||
|
||||
* [04_epilogue_visitor](/examples/python/04_epilogue_visitor.ipynb)
|
||||
|
||||
Shows how to fuse elementwise activation functions to GEMMs via the Python Epilogue Visitor interface
|
||||
|
||||
Reference in New Issue
Block a user