Updates for CUTLASS 3.5.0 (#1468)

This commit is contained in:
Vijay Thakkar
2024-04-11 21:33:40 -04:00
committed by GitHub
parent a40e08e9d5
commit 7d49e6c7e2
171 changed files with 7526 additions and 1888 deletions
@@ -1093,7 +1093,7 @@ private:
else if constexpr (ModeHasScales) {
Tensor sS = make_tensor(make_smem_ptr(shared_tensors.smem_scale.begin()), SmemLayoutScale{}); // (BLK_M,BLK_SCALE_K,PIPE)
Tensor tCsS = thread_mma.partition_A(sS);
Tensor tCrS = make_fragment_like<ElementScale>(thread_mma.partition_fragment_A(sS(_,_,Int<0>{})));
Tensor tCrS = make_tensor<ElementScale>(thread_mma.partition_fragment_A(sS(_,_,Int<0>{})).shape());
if constexpr (KernelConversionMode == ConversionMode::ConvertAndScale) {
return cute::make_tuple(tCsS, tCrS);
@@ -1101,7 +1101,7 @@ private:
else if constexpr (KernelConversionMode == ConversionMode::ConvertAndScaleWithZero) {
Tensor sZ = make_tensor(make_smem_ptr(shared_tensors.smem_zero.begin()), SmemLayoutScale{}); // (BLK_M,BLK_SCALE_K,PIPE)
Tensor tCsZ = thread_mma.partition_A(sZ);
Tensor tCrZ = make_fragment_like<ElementZero>(thread_mma.partition_fragment_A(sZ(_,_,Int<0>{})));
Tensor tCrZ = make_tensor<ElementZero>(thread_mma.partition_fragment_A(sZ(_,_,Int<0>{})).shape());
return cute::make_tuple(tCsS, tCrS, tCsZ, tCrZ);
}
else {
@@ -96,7 +96,7 @@ public:
using ElementB = typename GemmKernel::ElementB;
using ElementC = typename GemmKernel::ElementC;
using ElementD = typename GemmKernel::ElementD;
using ElementAccumulator = typename GemmKernel::TiledMma::ValTypeC;
using ElementAccumulator = typename GemmKernel::ElementAccumulator;
using DispatchPolicy = typename GemmKernel::DispatchPolicy;
using CollectiveMainloop = typename GemmKernel::CollectiveMainloop;
using CollectiveEpilogue = typename GemmKernel::CollectiveEpilogue;
@@ -361,9 +361,13 @@ public:
CUTLASS_ASSERT(cuda_adapter);
if (cuda_adapter) {
launch_result = cuda_adapter->launch(
grid, cluster, block, smem_size, stream, kernel_params, 0
);
launch_result = cuda_adapter->launch(grid,
cluster,
block,
smem_size,
stream,
kernel_params,
0);
}
else {
return Status::kErrorInternal;
@@ -32,16 +32,6 @@
\brief Defines common types used for all GEMM-like operators.
*/
/*
Note: CUTLASS 3x increases the host compiler requirements to C++17. However, certain
existing integrations of CUTLASS require C++11 host compilers.
Until this requirement can be lifted, certain headers with this annotation are required
to be remain consistent with C++11 syntax.
C++11 compatibility is enforced by `cutlass_test_unit_core_cpp11`.
*/
#pragma once
#include "cutlass/cutlass.h"
+23 -23
View File
@@ -30,7 +30,7 @@
**************************************************************************************************/
/*! \file
\brief
\brief
*/
#pragma once
@@ -177,8 +177,8 @@ public:
int const *ptr_scatter_D_indices = nullptr)
:
UniversalArgumentsBase(mode, problem_size, batch_count, batch_stride_D),
epilogue(epilogue),
ptr_A(ptr_A), ptr_B(ptr_B), ptr_C(ptr_C), ptr_D(ptr_D),
epilogue(epilogue),
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),
@@ -486,18 +486,18 @@ public:
int offset_k = 0;
int problem_size_k = params.problem_size.k();
ElementA *ptr_A = static_cast<ElementA *>(params.ptr_A);
ElementA *ptr_A = static_cast<ElementA *>(params.ptr_A);
ElementB *ptr_B = static_cast<ElementB *>(params.ptr_B);
//
// 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;
@@ -566,10 +566,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);
//
@@ -592,13 +592,13 @@ 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);
@@ -606,7 +606,7 @@ public:
// 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();
@@ -647,14 +647,14 @@ public:
);
Epilogue epilogue(
shared_storage.epilogue,
thread_idx,
warp_idx,
shared_storage.epilogue,
thread_idx,
warp_idx,
lane_idx);
// 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.
if (threadblock_tile_offset.k()) {
iterator_C = iterator_D;
@@ -666,11 +666,11 @@ public:
// Execute the epilogue operator to update the destination tensor.
epilogue(
output_op,
iterator_D,
accumulators,
iterator_C);
output_op,
iterator_D,
accumulators,
iterator_C);
//
// Release the semaphore
//
@@ -687,7 +687,7 @@ public:
// Otherwise, the semaphore is incremented
lock = threadblock_tile_offset.k() + 1;
}
semaphore.release(lock);
}
}
@@ -69,7 +69,6 @@ public:
using ProblemShape = ProblemShape_;
static_assert(cute::rank(ProblemShape{}) == 3 or cute::rank(ProblemShape{}) == 4,
"ProblemShape{} should be <M,N,K> or <M,N,K,L>");
// Mainloop derived types
using CollectiveMainloop = CollectiveMainloop_;
using TileShape = typename CollectiveMainloop::TileShape;
@@ -70,7 +70,6 @@ public:
using ProblemShape = ProblemShape_;
static_assert(cute::rank(ProblemShape{}) == 3 or cute::rank(ProblemShape{}) == 4,
"ProblemShape{} should be <M,N,K> or <M,N,K,L>");
// Mainloop derived types
using CollectiveMainloop = CollectiveMainloop_;
using TileShape = typename CollectiveMainloop::TileShape;
@@ -35,16 +35,6 @@
\brief Parameters structures for persistent tile schedulers
*/
/*
Note: CUTLASS 3x increases the host compiler requirements to C++17. However, certain
existing integrations of CUTLASS require C++11 host compilers.
Until this requirement can be lifted, certain headers with this annotation are required
to be remain consistent with C++11 syntax.
C++11 compatibility is enforced by this unit test: `cutlass_test_unit_core_cpp11`.
*/
#include "cutlass/coord.h"
#include "cutlass/kernel_hardware_info.h"
#include "cutlass/workspace.h"
+16 -33
View File
@@ -147,9 +147,7 @@ struct Mma_HFMA2 <
CUTLASS_PRAGMA_UNROLL
for(auto n=0; n < Shape::kN / Mma::Shape::kN; n++){
Array<half_t, 2> tmp;
Array<half_t, 2> *ptr_tmp = &tmp;
ptr_tmp[0] = ptr_D[n*Shape::kM/2 + m];
Array<half_t, 2> tmp { ptr_D[n*Shape::kM/2 + m] };
mma(
tmp,
@@ -157,7 +155,7 @@ struct Mma_HFMA2 <
ptr_B[n*Shape::kK + k],
tmp);
ptr_D[n*Shape::kM/2 + m] = ptr_tmp[0];
ptr_D[n*Shape::kM/2 + m] = tmp;
}
}
}
@@ -239,9 +237,7 @@ struct Mma_HFMA2<
CUTLASS_PRAGMA_UNROLL
for(auto m=0; m < Shape::kM / Mma::Shape::kM; m++){
Array<half_t, 2> tmp;
Array<half_t, 2> *ptr_tmp = &tmp;
ptr_tmp[0] = ptr_D[m*Shape::kN/2 + n];
Array<half_t, 2> tmp { ptr_D[m*Shape::kN/2 + n] };
Array<half_t, 2> tmp_B;
tmp_B[0] = ptr_B->at(2*n*Shape::kK + k);
@@ -253,7 +249,7 @@ struct Mma_HFMA2<
tmp_B,
tmp);
ptr_D[m*Shape::kN/2 + n] = ptr_tmp[0];
ptr_D[m*Shape::kN/2 + n] = tmp;
}
}
}
@@ -335,10 +331,7 @@ struct Mma_HFMA2 <
CUTLASS_PRAGMA_UNROLL
for (int n = 0; n < Shape::kN / Mma::Shape::kN; ++n) {
Array<half_t, 2> tmp;
Array<half_t, 2> *ptr_tmp = &tmp;
ptr_tmp[0] = ptr_D[m + n * Shape::kM/2];
Array<half_t, 2> tmp { ptr_D[m + n * Shape::kM/2] };
mma(
tmp,
@@ -346,7 +339,7 @@ struct Mma_HFMA2 <
ptr_B[k * Shape::kN + n],
tmp);
ptr_D[m + n * Shape::kM/2] = ptr_tmp[0];
ptr_D[m + n * Shape::kM/2] = tmp;
}
}
}
@@ -428,9 +421,7 @@ struct Mma_HFMA2<
CUTLASS_PRAGMA_UNROLL
for(auto m=0; m < Shape::kM / Mma::Shape::kM; m++){
Array<half_t, 2> tmp;
Array<half_t, 2> *ptr_tmp = &tmp;
ptr_tmp[0] = ptr_D[m*Shape::kN/2 + n];
Array<half_t, 2> tmp { ptr_D[m*Shape::kN/2 + n] };
mma(
tmp,
@@ -438,7 +429,7 @@ struct Mma_HFMA2<
ptr_B[k*Shape::kN/2 + n],
tmp);
ptr_D[m*Shape::kN/2 + n] = ptr_tmp[0];
ptr_D[m*Shape::kN/2 + n] = tmp;
}
}
}
@@ -521,9 +512,7 @@ struct Mma_HFMA2 <
CUTLASS_PRAGMA_UNROLL
for(auto n=0; n < Shape::kN / Mma::Shape::kN; n++){
Array<half_t, 2> tmp;
Array<half_t, 2> *ptr_tmp = &tmp;
ptr_tmp[0] = ptr_D[n*Shape::kM/2 + m];
Array<half_t, 2> tmp { ptr_D[n*Shape::kM/2 + m] };
Array<half_t, 2> tmp_A;
tmp_A[0] = ptr_A->at(2*m*Shape::kK + k);
@@ -535,7 +524,7 @@ struct Mma_HFMA2 <
ptr_B[n*Shape::kK + k],
tmp);
ptr_D[n*Shape::kM/2 + m] = ptr_tmp[0];
ptr_D[n*Shape::kM/2 + m] = tmp;
}
}
}
@@ -617,9 +606,7 @@ struct Mma_HFMA2 <
CUTLASS_PRAGMA_UNROLL
for(auto m=0; m < Shape::kM / Mma::Shape::kM; m++){
Array<half_t, 2> tmp;
Array<half_t, 2> *ptr_tmp = &tmp;
ptr_tmp[0] = ptr_D[m*Shape::kN/2 + n];
Array<half_t, 2> tmp { ptr_D[m*Shape::kN/2 + n] };
Array<half_t, 2> tmp_B;
tmp_B[0] = ptr_B->at(2*n*Shape::kK + k);
@@ -631,7 +618,7 @@ struct Mma_HFMA2 <
tmp_B,
tmp);
ptr_D[m*Shape::kN/2 + n] = ptr_tmp[0];
ptr_D[m*Shape::kN/2 + n] = tmp;
}
}
}
@@ -713,9 +700,7 @@ struct Mma_HFMA2 <
CUTLASS_PRAGMA_UNROLL
for(auto n=0; n < Shape::kN / Mma::Shape::kN; n++){
Array<half_t, 2> tmp;
Array<half_t, 2> *ptr_tmp = &tmp;
ptr_tmp[0] = ptr_D[n*Shape::kM/2 + m];
Array<half_t, 2> tmp { ptr_D[n*Shape::kM/2 + m] };
Array<half_t, 2> tmp_A;
tmp_A[0] = ptr_A->at(2*m*Shape::kK + k);
@@ -727,7 +712,7 @@ struct Mma_HFMA2 <
ptr_B[k*Shape::kN + n],
tmp);
ptr_D[n*Shape::kM/2 + m] = ptr_tmp[0];
ptr_D[n*Shape::kM/2 + m] = tmp;
}
}
}
@@ -810,9 +795,7 @@ struct Mma_HFMA2<
CUTLASS_PRAGMA_UNROLL
for(auto m=0; m < Shape::kM / Mma::Shape::kM; m++){
Array<half_t, 2> tmp;
Array<half_t, 2> *ptr_tmp = &tmp;
ptr_tmp[0] = ptr_D[m*Shape::kN/2 + n];
Array<half_t, 2> tmp { ptr_D[m*Shape::kN/2 + n] };
mma(
tmp,
@@ -820,7 +803,7 @@ struct Mma_HFMA2<
ptr_B[k*Shape::kN/2 + n],
tmp);
ptr_D[m*Shape::kN/2 + n] = ptr_tmp[0];
ptr_D[m*Shape::kN/2 + n] = tmp;
}
}
}
@@ -32,16 +32,6 @@
\brief Implements streamk threadblock mapping blockIdx to GEMM problems.
*/
/*
Note: CUTLASS 3x increases the host compiler requirements to C++17. However, certain
existing integrations of CUTLASS require C++11 host compilers.
Until this requirement can be lifted, certain headers with this annotation are required
to be remain consistent with C++11 syntax.
C++11 compatibility is enforced by `cutlass_test_unit_core_cpp11`.
*/
#pragma once
#include "cutlass/cutlass.h"