CUTLASS 3.0.0 (#786)

* CUTLASS 3.0.0
This commit is contained in:
Vijay Thakkar
2023-01-23 17:55:28 -08:00
committed by GitHub
parent 66d9cddc83
commit 277bd6e537
377 changed files with 76396 additions and 1186 deletions

View File

@@ -0,0 +1,54 @@
/***************************************************************************************************
* Copyright (c) 2017 - 2023 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.
*
**************************************************************************************************/
/* \file
\brief binding CuTe C++ APIs to Python
*/
#include <pybind11/pybind11.h>
#include <pybind11/stl_bind.h>
#include "cute/arch/mma_sm90_gmma.hpp"
namespace py = pybind11;
PYBIND11_MODULE(cute, m) {
// module doc
m.doc() = "CuTe C++ bindings";
py::enum_<cute::GMMA::Major>(m, "GMMAMajor",
R"pbdoc(classification of CuTe GMMA tensor major specification)pbdoc")
.value("K", cute::GMMA::Major::K,
R"pbdoc(Tensor is contiguous in reduction dimension)pbdoc")
.value("MN", cute::GMMA::Major::MN,
R"pbdoc(Tensor is contiguous in non-reduction dimension)pbdoc");
}

View File

@@ -29,8 +29,9 @@
*
**************************************************************************************************/
/* \file
\brief binding cutlass C++ APIs to python
\brief binding CUTLASS C++ APIs to Python
*/
#include <pybind11/pybind11.h>
#include <pybind11/stl_bind.h>

View File

@@ -34,6 +34,7 @@
\brief A generic wrapper around an epilogue visitor operation
*/
#pragma once
#include "cutlass/cutlass.h"

View File

@@ -30,8 +30,8 @@
**************************************************************************************************/
/*! \file
\brief Binary operations to be used within the epilogue visitor model.
\brief A file contains the binary ops
*/
#pragma once
@@ -44,7 +44,7 @@ namespace cutlass {
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Elementwise addition of two arrays
/// Scalar multiplication
template <typename T, int N>
struct VectorAdd {

View File

@@ -30,8 +30,8 @@
**************************************************************************************************/
/*! \file
\brief Unary operations to be used within the epilogue visitor model.
\brief A file contains the unary ops
*/
#pragma once

View File

@@ -30,8 +30,8 @@
**************************************************************************************************/
/*! \file
\brief Epilogue visitor operation that simply returns the accumulator
\brief A file contains the epilogue visitor Op with accumulator
*/
#pragma once

View File

@@ -30,8 +30,8 @@
**************************************************************************************************/
/*! \file
\brief Epilogue visitor operator performing a binary operation between two visitor nodes
\brief A file contains the epilogue visitor Op with Binary op
*/
#pragma once
@@ -84,7 +84,6 @@ public:
/// Fragment type of accumulator
using AccumulatorAccessType = Array<ElementAccumulator, kElementsPerAccess>;
/// Combination Op TODO: generalize this
using BinaryOp = BinaryOp_<ElementCompute, kElementsPerAccess>;
static_assert(kElementsPerAccess==VisitAccessTypeA::kElements, "kElementsPerAccess mismatches with Visitor A");

View File

@@ -30,8 +30,8 @@
**************************************************************************************************/
/*! \file
\brief Epilogue visitor operation that broadcasts a vector to all columns
\brief A file contains the epilogue visitor Op with broadcasting vector to all columns
*/
#pragma once

View File

@@ -30,8 +30,8 @@
**************************************************************************************************/
/*! \file
\brief Epilogue visitor operation that performs a column-wise reduction within a threadblock
\brief A file contains the epilogue visitor Op with reduction over columns in CTA
*/
#pragma once
@@ -68,7 +68,6 @@ public:
static int const kElementsPerAccess = OutputTileIterator::kElementsPerAccess;
// TODO: generalize the reduction op
using ReductionOp = cutlass::plus<Array<ElementReductionAccumulator, kElementsPerAccess>>;
using ReductionOpScalar = cutlass::plus<ElementReductionAccumulator>;
using ElementOutput = typename OutputTileIterator::Element;

View File

@@ -30,8 +30,8 @@
**************************************************************************************************/
/*! \file
\brief Epilogue visitor operation that performs a linear combination of two visitor nodes
\brief A file contains the epilogue visitor Op with Linear Combination
*/
#pragma once
@@ -82,7 +82,7 @@ public:
/// Fragment type of accumulator
using AccumulatorAccessType = Array<ElementAccumulator, kElementsPerAccess>;
/// Combination Op TODO: generalize this
/// Combination Op
using CombinationOp = cutlass::plus<VisitAccessType>;
static_assert(kElementsPerAccess==VisitAccessTypeA::kElements, "kElementsPerAccess mismatches with Visitor A");

View File

@@ -30,8 +30,8 @@
**************************************************************************************************/
/*! \file
\brief Epilogue visitor operation that broadcasts a vector to all rows
\brief A file contains the epilogue visitor Op with broadcasting vector to all rows
*/
#pragma once

View File

@@ -30,8 +30,8 @@
**************************************************************************************************/
/*! \file
\brief Epilogue visitor operation that performs a column-wise reduction within a threadblock
\brief A file contains the epilogue visitor Op with reduction over rows in CTA
*/
#pragma once
@@ -69,7 +69,6 @@ public:
static int const kElementsPerAccess = OutputTileIterator::kElementsPerAccess;
// TODO: generalize the reduction op
using ReductionOp = cutlass::plus<Array<ElementReductionAccumulator, kElementsPerAccess>>;
using ReductionOpScalar = cutlass::plus<ElementReductionAccumulator>;
using ElementOutput = typename OutputTileIterator::Element;

View File

@@ -30,8 +30,8 @@
**************************************************************************************************/
/*! \file
\brief Epilogue visitor operator performing a unary operation atop a visitor node
\brief A file contains the epilogue visitor Op with Unary operation
*/
#pragma once
@@ -79,7 +79,7 @@ public:
/// Fragment type of accumulator
using AccumulatorAccessType = Array<ElementAccumulator, kElementsPerAccess>;
/// Combination Op TODO: generalize this
/// Combination Op
using UnaryOp = UnaryOp_<ElementCompute, kElementsPerAccess>;
static_assert(kElementsPerAccess==VisitAccessTypeVisitor::kElements, "kElementsPerAccess mismatches with Visitor");

View File

@@ -30,7 +30,7 @@
**************************************************************************************************/
/*! \file
\brief
\brief
*/
#pragma once
@@ -139,8 +139,8 @@ public:
//
// Methods
//
Arguments():
Arguments():
ptr_A(nullptr), ptr_B(nullptr), ptr_C(nullptr), ptr_D(nullptr),
ptr_gather_A_indices(nullptr),
ptr_gather_B_indices(nullptr),
@@ -169,8 +169,8 @@ public:
int const *ptr_scatter_D_indices = nullptr
):
UniversalArgumentsBase(mode, problem_size, batch_count, batch_stride_D),
epilogue_visitor(epilogue_visitor),
ptr_A(ptr_A), ptr_B(ptr_B), ptr_C(ptr_C), ptr_D(ptr_D),
epilogue_visitor(epilogue_visitor),
ptr_A(ptr_A), ptr_B(ptr_B), ptr_C(ptr_C), ptr_D(ptr_D),
batch_stride_A(batch_stride_A), batch_stride_B(batch_stride_B), batch_stride_C(batch_stride_C),
stride_a(stride_a), stride_b(stride_b), stride_c(stride_c), stride_d(stride_d),
ptr_gather_A_indices(ptr_gather_A_indices), ptr_gather_B_indices(ptr_gather_B_indices),
@@ -205,8 +205,8 @@ public:
int const *ptr_scatter_D_indices = nullptr
):
UniversalArgumentsBase(mode, problem_size, batch_count, batch_stride_D),
epilogue_visitor(epilogue_visitor),
ptr_A(ptr_A), ptr_B(ptr_B), ptr_C(ptr_C), ptr_D(ptr_D),
epilogue_visitor(epilogue_visitor),
ptr_A(ptr_A), ptr_B(ptr_B), ptr_C(ptr_C), ptr_D(ptr_D),
batch_stride_A(batch_stride_A), batch_stride_B(batch_stride_B), batch_stride_C(batch_stride_C),
lda(lda), ldb(ldb), ldc(ldc), ldd(ldd),
ptr_gather_A_indices(ptr_gather_A_indices), ptr_gather_B_indices(ptr_gather_B_indices),
@@ -221,7 +221,7 @@ public:
/// Returns arguments for the transposed problem
Arguments transposed_problem() const {
Arguments args(*this);
std::swap(args.problem_size.m(), args.problem_size.n());
std::swap(args.ptr_A, args.ptr_B);
std::swap(args.lda, args.ldb);
@@ -256,7 +256,7 @@ public:
typename Mma::IteratorB::Params params_B;
typename EpilogueVisitor::OutputTileIterator::Params params_C;
typename EpilogueVisitor::OutputTileIterator::Params params_D;
typename EpilogueVisitor::Params epilogue_visitor;
void * ptr_A;
@@ -325,7 +325,7 @@ public:
batch_stride_C = args.batch_stride_C;
epilogue_visitor = args.epilogue_visitor;
semaphore = static_cast<int *>(workspace);
CUTLASS_TRACE_HOST("GemmUniversal::Params::update()");
}
@@ -345,7 +345,7 @@ public:
//
CUTLASS_DEVICE
GemmUniversalwithEpilogueVisitor() { }
GemmUniversalwithEpilogueVisitor() { }
/// Determines whether kernel satisfies alignment
static Status can_implement(
@@ -455,12 +455,12 @@ public:
//
// Fetch pointers based on mode.
//
if (params.mode == GemmUniversalMode::kGemm ||
if (params.mode == GemmUniversalMode::kGemm ||
params.mode == GemmUniversalMode::kGemmSplitKParallel) {
if (threadblock_tile_offset.k() + 1 < params.grid_tiled_shape.k()) {
problem_size_k = (threadblock_tile_offset.k() + 1) * params.gemm_k_size;
problem_size_k = (threadblock_tile_offset.k() + 1) * params.gemm_k_size;
}
offset_k = threadblock_tile_offset.k() * params.gemm_k_size;
@@ -529,10 +529,10 @@ public:
// Compute threadblock-scoped matrix multiply-add
mma(
gemm_k_iterations,
accumulators,
iterator_A,
iterator_B,
gemm_k_iterations,
accumulators,
iterator_A,
iterator_B,
accumulators);
//
@@ -555,30 +555,16 @@ public:
int block_idx = threadblock_tile_offset.m() + threadblock_tile_offset.n() * params.grid_tiled_shape.m();
ElementC *ptr_C = static_cast<ElementC *>(params.ptr_C);
ElementC *ptr_C = static_cast<ElementC *>(params.ptr_C);
ElementC *ptr_D = static_cast<ElementC *>(params.ptr_D);
//
// Fetch pointers based on mode.
//
// Construct the semaphore.
Semaphore semaphore(params.semaphore + block_idx, thread_idx);
// if (params.mode == GemmUniversalMode::kGemm) {
// // TODO: fix this order
// // If performing a reduction via split-K, fetch the initial synchronization
// if (params.grid_tiled_shape.k() > 1) {
// // Fetch the synchronization lock initially but do not block.
// semaphore.fetch();
// // Indicate which position in a serial reduction the output operator is currently updating
// output_op.set_k_partition(threadblock_tile_offset.k(), params.grid_tiled_shape.k());
// }
// }
// Tile iterator loading from source tensor.
EpilogueVisitor epilogue_visitor(
@@ -590,9 +576,6 @@ public:
params.problem_size.mn()
);
// if (params.mode == GemmUniversalMode::kGemmSplitKParallel) {
// ptr_D += threadblock_tile_offset.k() * params.batch_stride_D;
// }
if (params.mode == GemmUniversalMode::kBatched || params.mode == GemmUniversalMode::kArray) {
epilogue_visitor.set_batch_index(threadblock_tile_offset.k());
}
@@ -605,25 +588,20 @@ public:
// Wait on the semaphore - this latency may have been covered by iterator construction
if (params.mode == GemmUniversalMode::kGemm && params.grid_tiled_shape.k() > 1) {
// For subsequent threadblocks, the source matrix is held in the 'D' tensor.
// TODO: ???
// if (threadblock_tile_offset.k()) {
// iterator_C = iterator_D;
// }
// For subsequent threadblocks, the source matrix is held in the 'D' tensor.
semaphore.wait(threadblock_tile_offset.k());
}
// Execute the epilogue operator to update the destination tensor.
epilogue(epilogue_visitor, accumulators);
epilogue(epilogue_visitor, accumulators);
//
// Release the semaphore
//
if (params.mode == GemmUniversalMode::kGemm && params.grid_tiled_shape.k() > 1) {
if (params.mode == GemmUniversalMode::kGemm && params.grid_tiled_shape.k() > 1) {
int lock = 0;
if (params.grid_tiled_shape.k() == threadblock_tile_offset.k() + 1) {
@@ -635,7 +613,7 @@ public:
// Otherwise, the semaphore is incremented
lock = threadblock_tile_offset.k() + 1;
}
semaphore.release(lock);
}
}

View File

@@ -83,7 +83,6 @@ void bind_identity_swizzle(py::module & m, std::string name) {
:param problem_size: Implicit gemm problem size conv_operator(NZPQK, NDHWC, KTRSC)
:type problem_size: :class:`cutlass.gemm.GemmCoord`)
)pbdoc")
// TODO: the returned dim3 is not usable in python
.def("get_grid_shape", &T::get_grid_shape,
py::arg("tiled_shape"),
R"pbdoc(Computes CUDA grid dimensions given a size in units of logical tiles)pbdoc")