@@ -246,6 +246,9 @@ compute_lower_srt(ConvProblemShape<ConvOp, NumSpatialDimensions> const& problem_
|
||||
return lower;
|
||||
}
|
||||
|
||||
template <class CopyOp> struct is_im2col_load { static constexpr bool value = false; };
|
||||
template <> struct is_im2col_load<SM90_TMA_LOAD_IM2COL > { static constexpr bool value = true; };
|
||||
template <> struct is_im2col_load<SM90_TMA_LOAD_IM2COL_MULTICAST> { static constexpr bool value = true; };
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::conv::collective::detail
|
||||
|
||||
@@ -131,6 +131,9 @@ struct CollectiveConv<
|
||||
&& (cute::is_same_v<GmemTiledCopyB, SM90_TMA_LOAD> || cute::is_same_v<GmemTiledCopyB, SM90_TMA_LOAD_MULTICAST>)),
|
||||
"GmemTiledCopyB - invalid SM90 TMA copy atom specified.");
|
||||
|
||||
static constexpr bool is_im2col_A = detail::is_im2col_load<GmemTiledCopyA>::value;
|
||||
static constexpr bool is_im2col_B = detail::is_im2col_load<GmemTiledCopyB>::value;
|
||||
|
||||
// TMA converts f32 input to tf32 when copying from GMEM to SMEM
|
||||
// For all other types, cast to size equivalent uint type to avoid any rounding by TMA.
|
||||
static constexpr bool ConvertF32toTF32A = cute::is_same_v<float, ElementA>;
|
||||
@@ -169,12 +172,11 @@ private:
|
||||
// Note that for fprop and dgrad kernel, the tma load mode is im2col for tensor A and tiled for
|
||||
// tensor B while for wgrad kernel, the tma load mode is tiled for tensor A and im2col for tensor
|
||||
// B since operand A, B is swapped.
|
||||
|
||||
// Get tma_load_a instantce.
|
||||
template <class TensorA>
|
||||
static constexpr auto
|
||||
get_tma_load_a_instance(TensorA const& tensor_a, typename Arguments::ProblemShape const& problem_shape) {
|
||||
if constexpr (ConvOp == conv::Operator::kFprop || ConvOp == conv::Operator::kDgrad) {
|
||||
if constexpr (is_im2col_A) {
|
||||
// compute the upper and lower corners based on the conv padding
|
||||
auto lower_corner_whd = detail::compute_lower_corner_whd(problem_shape);
|
||||
auto upper_corner_whd = detail::compute_upper_corner_whd(problem_shape);
|
||||
@@ -203,7 +205,7 @@ private:
|
||||
shape(stride_srt));
|
||||
}
|
||||
// TMA tiled mode for tensor A in wgrad kernel.
|
||||
else if constexpr (ConvOp == conv::Operator::kWgrad) {
|
||||
else {
|
||||
return make_tma_copy(
|
||||
GmemTiledCopyA{},
|
||||
tensor_a,
|
||||
@@ -217,16 +219,8 @@ private:
|
||||
template <class TensorB>
|
||||
static constexpr auto
|
||||
get_tma_load_b_instance(TensorB const& tensor_b, typename Arguments::ProblemShape const& problem_shape) {
|
||||
if constexpr (ConvOp == conv::Operator::kFprop || ConvOp == conv::Operator::kDgrad) {
|
||||
return make_tma_copy(
|
||||
GmemTiledCopyB{},
|
||||
tensor_b,
|
||||
SmemLayoutB{}(_,_,_0{}),
|
||||
make_shape(shape<1>(TileShape{}), shape<2>(TileShape{})),
|
||||
size<0>(ClusterShape{}));
|
||||
}
|
||||
// TMA im2col mode for tensor B in wgrad kernel.
|
||||
else if constexpr (ConvOp == conv::Operator::kWgrad) {
|
||||
if constexpr (is_im2col_B) {
|
||||
// compute the upper and lower corners based on the conv padding
|
||||
auto lower_corner_whd = detail::compute_lower_corner_whd(problem_shape);
|
||||
auto upper_corner_whd = detail::compute_upper_corner_whd(problem_shape);
|
||||
@@ -246,6 +240,26 @@ private:
|
||||
shape(lower_srt),
|
||||
cute::reverse(shape(problem_shape.dilation)));
|
||||
}
|
||||
else {
|
||||
return make_tma_copy(
|
||||
GmemTiledCopyB{},
|
||||
tensor_b,
|
||||
SmemLayoutB{}(_,_,_0{}),
|
||||
make_shape(shape<1>(TileShape{}), shape<2>(TileShape{})),
|
||||
size<0>(ClusterShape{}));
|
||||
}
|
||||
}
|
||||
|
||||
static constexpr auto
|
||||
get_problem_shape_MNKL(typename Arguments::ProblemShape const& problem_shape) {
|
||||
if constexpr (is_im2col_A || is_im2col_B) {
|
||||
// transformation + im2col linearization
|
||||
return problem_shape.get_linearized_problem_shape_MNKL();
|
||||
}
|
||||
else {
|
||||
// transformation
|
||||
return problem_shape.get_transformed_problem_shape_MNKL();
|
||||
}
|
||||
}
|
||||
|
||||
public:
|
||||
@@ -253,9 +267,7 @@ public:
|
||||
// Device side kernel params
|
||||
struct Params {
|
||||
using _Submode = decltype(take<0,NumTensorDimensions-1>(typename Arguments::ProblemShape::TensorExtent{}));
|
||||
using ProblemShape = cute::conditional_t<DispatchPolicy::ConvOp == conv::Operator::kWgrad,
|
||||
Shape<int, _Submode, _Submode>,
|
||||
Shape<_Submode, int, _Submode>>;
|
||||
using ProblemShape = decltype(get_problem_shape_MNKL(typename Arguments::ProblemShape{}));
|
||||
|
||||
// Assumption: StrideA is congruent with Problem_MK
|
||||
// Select TMA load type according to convolution operator.
|
||||
@@ -283,6 +295,7 @@ public:
|
||||
TMA_A tma_load_a;
|
||||
TMA_B tma_load_b;
|
||||
ProblemShape problem_shape;
|
||||
uint32_t tma_transaction_bytes = TmaTransactionBytes;
|
||||
};
|
||||
|
||||
//
|
||||
@@ -314,17 +327,18 @@ public:
|
||||
auto tma_load_a = get_tma_load_a_instance(tensor_a, args.problem_shape);
|
||||
auto tma_load_b = get_tma_load_b_instance(tensor_b, args.problem_shape);
|
||||
|
||||
auto problem_shape_mnk = args.problem_shape.get_transformed_problem_shape_MNK();
|
||||
auto problem_shape_mnkl = get_problem_shape_MNKL(args.problem_shape);
|
||||
|
||||
return {
|
||||
tma_load_a,
|
||||
tma_load_b,
|
||||
problem_shape_mnk
|
||||
problem_shape_mnkl,
|
||||
TmaTransactionBytes
|
||||
};
|
||||
}
|
||||
|
||||
template<class ProblemShape>
|
||||
CUTLASS_HOST_DEVICE static bool
|
||||
static bool
|
||||
can_implement(
|
||||
ProblemShape const& problem_shape,
|
||||
Arguments const& args) {
|
||||
@@ -389,13 +403,12 @@ public:
|
||||
TensorA const& gA, TMA_LOAD_A& tma_load_a,
|
||||
TensorB const& gB, TMA_LOAD_B& tma_load_b,
|
||||
KTileIterator k_tile_iter, int k_tile_count,
|
||||
int thad_idx,
|
||||
int thread_idx,
|
||||
uint32_t block_rank_in_cluster,
|
||||
TensorStorage& shared_tensors) {
|
||||
int warp_idx = canonical_warp_idx_sync();
|
||||
int warp_idx_in_warp_group = warp_idx % 4;
|
||||
int lane_predicate = cute::elect_one_sync();
|
||||
|
||||
if (warp_idx_in_warp_group == 0 and lane_predicate) {
|
||||
if (lane_predicate) {
|
||||
Tensor sA = make_tensor(make_smem_ptr(shared_tensors.smem_A.data()), SmemLayoutA{}); // (BLK_M,BLK_K,PIPE)
|
||||
Tensor sB = make_tensor(make_smem_ptr(shared_tensors.smem_B.data()), SmemLayoutB{}); // (BLK_N,BLK_K,PIPE)
|
||||
|
||||
@@ -403,7 +416,8 @@ public:
|
||||
// Prepare the TMA loads for A and B
|
||||
//
|
||||
|
||||
dim3 cluster_local_block_id = cute::block_id_in_cluster();
|
||||
constexpr uint32_t cluster_shape_x = get<0>(ClusterShape());
|
||||
uint2 cluster_local_block_id = {block_rank_in_cluster % cluster_shape_x, block_rank_in_cluster / cluster_shape_x};
|
||||
auto block_tma_a = tma_load_a.get_slice(cluster_local_block_id.y);
|
||||
auto block_tma_b = tma_load_b.get_slice(cluster_local_block_id.x);
|
||||
|
||||
@@ -462,12 +476,10 @@ public:
|
||||
/// Perform a Producer Epilogue to prevent early exit of blocks in a Cluster
|
||||
CUTLASS_DEVICE void
|
||||
load_tail(MainloopPipeline pipeline, PipelineState smem_pipe_producer_state) {
|
||||
int warp_idx = canonical_warp_idx_sync();
|
||||
int warp_idx_in_warp_group = warp_idx % 4;
|
||||
int lane_predicate = cute::elect_one_sync();
|
||||
|
||||
// Issue the epilogue waits
|
||||
if (warp_idx_in_warp_group == 0 and lane_predicate) {
|
||||
if (lane_predicate) {
|
||||
/* This helps avoid early exit of blocks in Cluster
|
||||
* Waits for all stages to either be released (all
|
||||
* Consumer UNLOCKs), or if the stage was never used
|
||||
|
||||
@@ -352,15 +352,16 @@ struct ConvProblemShape {
|
||||
}
|
||||
}
|
||||
|
||||
// Get problem shape MNK according to following table:
|
||||
// | | Fprop | Dgrad | Wgrad |
|
||||
// | ---- | --------- | -------- | -------- |
|
||||
// | Shape_M | (Q,P,Z,N) | (W,H,D,N) | (K) |
|
||||
// | Shape_N | (K) | (C) | (C,S,R,T) |
|
||||
// | Shape_K | (C,S,R,T) | (K,S,R,T) | (Q,P,Z,N) |
|
||||
// Get problem shape MNKL according to following table:
|
||||
// | | Fprop | Dgrad | Wgrad |
|
||||
// | ---- | --------- | -------- | -------- |
|
||||
// | Shape_M | (Q,P,Z,N) | (W/V,H/U,D/O,N) | (K) |
|
||||
// | Shape_N | (K) | (C) | (C,S,R,T) |
|
||||
// | Shape_K | (C,S,R,T) | (K,S,R,T) | (Q,P,Z,N) |
|
||||
// | Shape_L | _1 | (V,U,O) | _1 |
|
||||
CUTLASS_HOST_DEVICE
|
||||
constexpr auto
|
||||
get_transformed_problem_shape_MNK() const {
|
||||
get_transformed_problem_shape_MNKL() const {
|
||||
using cute::insert;
|
||||
using cute::make_shape;
|
||||
using cute::reverse;
|
||||
@@ -370,32 +371,56 @@ struct ConvProblemShape {
|
||||
auto M_xformed = shape_C[0];
|
||||
auto N_xformed = reverse(take<1, RankT>(shape_C));
|
||||
auto K_xformed = reverse(take<0, RankT - 1>(shape_A));
|
||||
auto L_xformed = cute::Int<1>{};
|
||||
|
||||
return make_shape(M_xformed, N_xformed, K_xformed);
|
||||
return make_shape(M_xformed, N_xformed, K_xformed, L_xformed);
|
||||
}
|
||||
else if constexpr (ConvOp == conv::Operator::kFprop){
|
||||
auto M_xformed = reverse(take<0, RankT - 1>(shape_C));
|
||||
auto N_xformed = shape_C[RankT - 1];
|
||||
auto K_xformed = reverse(take<1, RankT>(shape_B));
|
||||
auto L_xformed = cute::Int<1>{};
|
||||
|
||||
return make_shape(M_xformed, N_xformed, K_xformed);
|
||||
return make_shape(M_xformed, N_xformed, K_xformed, L_xformed);
|
||||
}
|
||||
else if constexpr (ConvOp == conv::Operator::kDgrad) {
|
||||
auto M_xformed = reverse(take<0,RankT - 1>(shape_C));
|
||||
auto L_xformed = reverse(traversal_stride); // (V,U,O)
|
||||
auto M_xformed = ceil_div(reverse(take<0,RankT - 1>(shape_C)), L_xformed);
|
||||
auto N_xformed = shape_C[RankT - 1];
|
||||
// shape_B: [K,T,R,S,C], K_xformed: [K,S,R,T]
|
||||
auto K_xformed = insert<0>(
|
||||
(reverse(take<1,RankT - 1>(shape_B))),
|
||||
shape_B[0]);
|
||||
return make_shape(M_xformed, N_xformed, K_xformed);
|
||||
|
||||
return make_shape(M_xformed, N_xformed, K_xformed, L_xformed);
|
||||
}
|
||||
}
|
||||
|
||||
// Assuming im2col linearization
|
||||
// Get problem shape MNKL according to following table:
|
||||
// | | Fprop | Dgrad | Wgrad |
|
||||
// | ---- | --------- | -------- | -------- |
|
||||
// | Shape_M | (Q*P*Z*N) | ([W/V]*[H/U]*[D/O]*N) | (K) |
|
||||
// | Shape_N | (K) | (C) | (C,S,R,T) |
|
||||
// | Shape_K | (C,S,R,T) | (K,S,R,T) | (Q*P*Z*N) |
|
||||
// | Shape_L | _1 | (V*U*O) | _1 |
|
||||
CUTLASS_HOST_DEVICE
|
||||
constexpr auto
|
||||
get_linearized_problem_shape_MNKL() const {
|
||||
auto [M, N, K, L] = get_transformed_problem_shape_MNKL();
|
||||
|
||||
if constexpr (ConvOp == conv::Operator::kFprop || ConvOp == conv::Operator::kDgrad) {
|
||||
return cute::make_shape(cute::product(M), N, K, cute::product(L));
|
||||
}
|
||||
else if constexpr (ConvOp == conv::Operator::kWgrad) {
|
||||
return cute::make_shape(M, N, cute::product(K), L);
|
||||
}
|
||||
}
|
||||
|
||||
// Get A extents.
|
||||
// fprop: A extents array contains [N,D,H,W,C]. Turn that into ((W,H,D,N), (C))
|
||||
// wgrad: A extents array contains [N,Z,P,Q,K]. Turn that into ((K), (Q,P,Z,N))
|
||||
// dgrad: A extents array contains [N,Z,P,Q,K]. Turn that into ((Q,P,Z,N), (K))
|
||||
// wgrad: A extents array contains [N,Z,P,Q,K]. Turn that into ((K), (Q,P,Z,N))
|
||||
CUTLASS_HOST_DEVICE
|
||||
constexpr auto
|
||||
get_shape_A() const {
|
||||
@@ -418,8 +443,8 @@ struct ConvProblemShape {
|
||||
|
||||
// Get B extents.
|
||||
// fprop: B extents array contains [K,T,R,S,C]. Turn that into ((K), (C,S,R,T))
|
||||
// wgrad: B extents array contains [N,D,H,W,C]. Turn that into ((C), (W,H,D,N))
|
||||
// dgrad: B extents array contains [K,T,R,S,C]. Turn that into ((C), (K,S,R,T))
|
||||
// wgrad: B extents array contains [N,D,H,W,C]. Turn that into ((C), (W,H,D,N))
|
||||
CUTLASS_HOST_DEVICE
|
||||
constexpr auto
|
||||
get_shape_B() const {
|
||||
@@ -447,6 +472,30 @@ struct ConvProblemShape {
|
||||
}
|
||||
}
|
||||
|
||||
// Get C extents.
|
||||
// fprop: C extents array contains [N,Z,P,Q,K]. Turn that into ((Q,P,Z,N), (K))
|
||||
// dgrad: C extents array contains [N,D,H,W,C]. Turn that into ((W,H,D,N), (C))
|
||||
// wgrad: C extents array contains [K,T,R,S,C]. Turn that into ((K), (C,S,R,T))
|
||||
CUTLASS_HOST_DEVICE
|
||||
constexpr auto
|
||||
get_shape_C() const {
|
||||
using cute::make_shape;
|
||||
using cute::reverse;
|
||||
using cute::take;
|
||||
|
||||
if constexpr (ConvOp == conv::Operator::kFprop ||
|
||||
ConvOp == conv::Operator::kDgrad) {
|
||||
return make_shape(
|
||||
reverse(take<0, RankT - 1>(shape_C)),
|
||||
shape_C[RankT - 1]);
|
||||
}
|
||||
else if constexpr (ConvOp == conv::Operator::kWgrad) {
|
||||
return make_shape(
|
||||
shape_C[0],
|
||||
reverse(take<1, RankT>(shape_C)));
|
||||
}
|
||||
}
|
||||
|
||||
// Static method that returns the canonical strides of tensors (layouts are right major and compact)
|
||||
CUTLASS_HOST_DEVICE
|
||||
static constexpr TensorStride
|
||||
@@ -529,7 +578,9 @@ private:
|
||||
// calculate n,z,p,q,k.
|
||||
// a helper lambda to compute a single spatial extent of the nzpqk tensor
|
||||
auto nzpqk_extent = [](int act_ext, int filter_ext, int pad_total, int dilation, int tstride) {
|
||||
return 1 + (act_ext + pad_total - ((filter_ext -1) * dilation + 1)) / tstride;
|
||||
auto tmp = act_ext + pad_total - ((filter_ext -1) * dilation + 1);
|
||||
CUTLASS_ASSERT(tmp % tstride == 0);
|
||||
return 1 + tmp / tstride;
|
||||
};
|
||||
|
||||
shape_xformed_act[0] = shape_act[0]; // Activation N extent
|
||||
|
||||
@@ -228,29 +228,18 @@ public:
|
||||
/// Initializes conv state from arguments.
|
||||
Status
|
||||
initialize(
|
||||
Arguments const& args,
|
||||
void* workspace = nullptr,
|
||||
cudaStream_t stream = nullptr,
|
||||
CudaHostAdapter *cuda_adapter = nullptr) {
|
||||
Arguments const& args,
|
||||
void* workspace = nullptr,
|
||||
cudaStream_t stream = nullptr,
|
||||
CudaHostAdapter *cuda_adapter = nullptr) {
|
||||
|
||||
CUTLASS_TRACE_HOST("ConvUniversal::initialize() - workspace "
|
||||
<< workspace << ", stream: " << (stream ? "non-null" : "null"));
|
||||
|
||||
size_t workspace_bytes = ConvKernel::get_workspace_size(args);
|
||||
CUTLASS_TRACE_HOST(" workspace_bytes: " << workspace_bytes);
|
||||
|
||||
if (workspace_bytes) {
|
||||
if (!workspace) {
|
||||
CUTLASS_TRACE_HOST(" error: device workspace must not be null");
|
||||
return Status::kErrorWorkspaceNull;
|
||||
}
|
||||
|
||||
CUTLASS_TRACE_HOST(" clearing device workspace");
|
||||
cudaError_t result = cudaMemsetAsync(workspace, 0, workspace_bytes, stream);
|
||||
if (cudaSuccess != result) {
|
||||
result = cudaGetLastError(); // to clear the error bit
|
||||
CUTLASS_TRACE_HOST(" cudaMemsetAsync() returned error " << cudaGetErrorString(result));
|
||||
return Status::kErrorInternal;
|
||||
}
|
||||
// Initialize the workspace
|
||||
Status status = ConvKernel::initialize_workspace(args, workspace, stream, cuda_adapter);
|
||||
if (status != Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
// Initialize the Params structure
|
||||
@@ -297,7 +286,7 @@ public:
|
||||
/// Primary run() entry point API that is static allowing users to create and manage their own params.
|
||||
/// Supplied params struct must be construct by calling ConvKernel::to_underling_arguments()
|
||||
static Status
|
||||
run(Params& params, cudaStream_t stream = nullptr, CudaHostAdapter *cuda_adapter = nullptr) {
|
||||
run(Params& params, cudaStream_t stream = nullptr, CudaHostAdapter *cuda_adapter = nullptr, int32_t kernel_index = 0) {
|
||||
CUTLASS_TRACE_HOST("ConvUniversal::run()");
|
||||
dim3 const block = ConvKernel::get_block_shape();
|
||||
dim3 const grid = get_grid_shape(params);
|
||||
@@ -319,9 +308,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,
|
||||
kernel_index);
|
||||
}
|
||||
else {
|
||||
return Status::kErrorInternal;
|
||||
@@ -379,11 +372,12 @@ public:
|
||||
Arguments const& args,
|
||||
void* workspace = nullptr,
|
||||
cudaStream_t stream = nullptr,
|
||||
CudaHostAdapter *cuda_adapter = nullptr
|
||||
CudaHostAdapter *cuda_adapter = nullptr,
|
||||
int32_t kernel_index = 0
|
||||
) {
|
||||
Status status = initialize(args, workspace, stream, cuda_adapter);
|
||||
if (Status::kSuccess == status) {
|
||||
status = run(params_, stream, cuda_adapter);
|
||||
status = run(params_, stream, cuda_adapter, kernel_index);
|
||||
}
|
||||
return status;
|
||||
}
|
||||
|
||||
@@ -197,7 +197,7 @@ public:
|
||||
params_.ptr_C = args.ref_C.data();
|
||||
params_.ptr_D = args.ref_D.data();
|
||||
params_.output_op = args.output_op;
|
||||
params_.ptr_reordered_B = args.ref_reordered_B.data();;
|
||||
params_.ptr_reordered_B = args.ref_reordered_B.data();
|
||||
params_.semaphore = static_cast<int *>(workspace);
|
||||
|
||||
return Status::kSuccess;
|
||||
|
||||
@@ -31,6 +31,7 @@
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/conv/convolution.h"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "cutlass/arch/arch.h"
|
||||
|
||||
#include "cute/layout.hpp"
|
||||
@@ -38,6 +39,8 @@
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::conv {
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -114,7 +114,10 @@ template <
|
||||
typename ElementTensor,
|
||||
typename ElementVector,
|
||||
typename OutputOp,
|
||||
int ElementsPerAccess
|
||||
int ElementsPerAccess,
|
||||
typename PermuteDLayout = layout::NoPermute,
|
||||
conv::StrideSupport StrideSupport = conv::StrideSupport::kUnity,
|
||||
int Rank = 4
|
||||
>
|
||||
struct DefaultConvEpilogueWithBroadcastSimt {
|
||||
using Epilogue = typename epilogue::threadblock::DefaultEpilogueWithBroadcastSimt<
|
||||
@@ -124,7 +127,11 @@ struct DefaultConvEpilogueWithBroadcastSimt {
|
||||
ElementTensor,
|
||||
ElementVector,
|
||||
OutputOp,
|
||||
ElementsPerAccess
|
||||
ElementsPerAccess,
|
||||
false,
|
||||
PermuteDLayout,
|
||||
StrideSupport,
|
||||
Rank
|
||||
>::Epilogue;
|
||||
};
|
||||
|
||||
|
||||
@@ -197,7 +197,10 @@ struct DefaultConv3dFpropWithBroadcast <
|
||||
typename EpilogueOutputOp::ElementT,
|
||||
typename EpilogueOutputOp::ElementVector,
|
||||
EpilogueOutputOp,
|
||||
ImplicitGemmBase::Epilogue::kElementsPerAccess
|
||||
ImplicitGemmBase::Epilogue::kElementsPerAccess,
|
||||
layout::NoPermute,
|
||||
StrideSupport,
|
||||
5
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
|
||||
@@ -181,7 +181,11 @@ struct DefaultDeconv2d <
|
||||
ThreadblockShape,
|
||||
WarpMmaSimtOp,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
EpilogueOutputOp::kCount,
|
||||
false,
|
||||
layout::NoPermute,
|
||||
StrideSupport::kStrided,
|
||||
4
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
@@ -405,7 +409,11 @@ struct DefaultDeconv2d <
|
||||
ThreadblockShape,
|
||||
WarpMmaSimtOp,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
EpilogueOutputOp::kCount,
|
||||
false,
|
||||
layout::NoPermute,
|
||||
StrideSupport::kStrided,
|
||||
4
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
@@ -627,7 +635,11 @@ struct DefaultDeconv2d <
|
||||
ThreadblockShape,
|
||||
WarpMmaSimtOp,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
EpilogueOutputOp::kCount,
|
||||
false,
|
||||
layout::NoPermute,
|
||||
StrideSupport::kStrided,
|
||||
4
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
@@ -852,7 +864,11 @@ struct DefaultDeconv2d <
|
||||
ThreadblockShape,
|
||||
WarpMmaSimtOp,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
EpilogueOutputOp::kCount,
|
||||
false,
|
||||
layout::NoPermute,
|
||||
StrideSupport::kStrided,
|
||||
4
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
|
||||
@@ -170,7 +170,11 @@ struct DefaultDeconv3d <
|
||||
ThreadblockShape,
|
||||
WarpMmaSimtOp,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
EpilogueOutputOp::kCount,
|
||||
false,
|
||||
layout::NoPermute,
|
||||
StrideSupport::kStrided,
|
||||
5
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
@@ -282,7 +286,11 @@ struct DefaultDeconv3d <
|
||||
ThreadblockShape,
|
||||
WarpMmaSimtOp,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
EpilogueOutputOp::kCount,
|
||||
false,
|
||||
layout::NoPermute,
|
||||
StrideSupport::kStrided,
|
||||
5
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
@@ -389,7 +397,11 @@ struct DefaultDeconv3d <
|
||||
ThreadblockShape,
|
||||
WarpMmaSimtOp,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
EpilogueOutputOp::kCount,
|
||||
false,
|
||||
layout::NoPermute,
|
||||
StrideSupport::kStrided,
|
||||
5
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
@@ -501,7 +513,11 @@ struct DefaultDeconv3d <
|
||||
ThreadblockShape,
|
||||
WarpMmaSimtOp,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
EpilogueOutputOp::kCount,
|
||||
false,
|
||||
layout::NoPermute,
|
||||
StrideSupport::kStrided,
|
||||
5
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
|
||||
@@ -196,7 +196,10 @@ struct DefaultDeconv3dWithBroadcast <
|
||||
typename EpilogueOutputOp::ElementT,
|
||||
typename EpilogueOutputOp::ElementVector,
|
||||
EpilogueOutputOp,
|
||||
ImplicitGemmBase::Epilogue::kElementsPerAccess
|
||||
ImplicitGemmBase::Epilogue::kElementsPerAccess,
|
||||
layout::NoPermute,
|
||||
StrideSupport::kStrided,
|
||||
5
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
@@ -273,7 +276,7 @@ struct DefaultDeconv3dWithBroadcast <
|
||||
>::Kernel;
|
||||
|
||||
// Define epilogue
|
||||
using Epilogue = typename cutlass::conv::kernel::detail::DefaultConvEpilogueWithBroadcastSimtStridedDgrad<
|
||||
using Epilogue = typename cutlass::conv::kernel::detail::DefaultConvEpilogueWithBroadcastSimt<
|
||||
ArchTag,
|
||||
typename ImplicitGemmBase::Epilogue::Shape,
|
||||
typename ImplicitGemmBase::Epilogue::WarpMmaOperator,
|
||||
@@ -281,7 +284,10 @@ struct DefaultDeconv3dWithBroadcast <
|
||||
typename EpilogueOutputOp::ElementT,
|
||||
typename EpilogueOutputOp::ElementVector,
|
||||
EpilogueOutputOp,
|
||||
ImplicitGemmBase::Epilogue::kElementsPerAccess
|
||||
ImplicitGemmBase::Epilogue::kElementsPerAccess,
|
||||
layout::NoPermute,
|
||||
StrideSupport::kStrided,
|
||||
5
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
|
||||
@@ -233,9 +233,9 @@ struct ImplicitGemmConvolution {
|
||||
ptr_A(args.ref_A.data()),
|
||||
iterator_B(args.problem_size, args.ref_B.layout()),
|
||||
ptr_B(args.ref_B.data()),
|
||||
iterator_C(ConvOutputIteratorParameter::layout(args.ref_C), args.problem_size),
|
||||
iterator_C(ConvOutputIteratorParameter::layout(args.ref_C), implicit_gemm_tensor_c_extent(kConvolutionalOperator, args.problem_size)),
|
||||
ptr_C(args.ref_C.data()),
|
||||
iterator_D(ConvOutputIteratorParameter::layout(args.ref_D), args.problem_size),
|
||||
iterator_D(ConvOutputIteratorParameter::layout(args.ref_D), implicit_gemm_tensor_c_extent(kConvolutionalOperator, args.problem_size)),
|
||||
ptr_D(args.ref_D.data()),
|
||||
output_op(args.output_op),
|
||||
semaphore(semaphore),
|
||||
|
||||
@@ -257,9 +257,9 @@ struct ImplicitGemmConvolutionWithFusedEpilogue {
|
||||
ptr_A(args.ref_A.data()),
|
||||
iterator_B(args.problem_size, args.ref_B.layout()),
|
||||
ptr_B(args.ref_B.data()),
|
||||
iterator_C(ConvOutputIteratorParameter::layout(args.ref_C)),
|
||||
iterator_C(ConvOutputIteratorParameter::layout(args.ref_C), implicit_gemm_tensor_c_extent(kConvolutionalOperator, args.problem_size)),
|
||||
ptr_C(args.ref_C.data()),
|
||||
iterator_D(ConvOutputIteratorParameter::layout(args.ref_D)),
|
||||
iterator_D(ConvOutputIteratorParameter::layout(args.ref_D), implicit_gemm_tensor_c_extent(kConvolutionalOperator, args.problem_size)),
|
||||
ptr_D(args.ref_D.data()),
|
||||
output_op(args.output_op),
|
||||
semaphore(semaphore),
|
||||
|
||||
@@ -51,12 +51,12 @@ namespace cutlass::conv::kernel {
|
||||
template <
|
||||
class CollectiveMainloop_,
|
||||
class CollectiveEpilogue_,
|
||||
class TileSchedulerTag
|
||||
class TileSchedulerTag_
|
||||
>
|
||||
class ConvUniversal<
|
||||
CollectiveMainloop_,
|
||||
CollectiveEpilogue_,
|
||||
TileSchedulerTag,
|
||||
TileSchedulerTag_,
|
||||
cute::enable_if_t<cute::is_base_of_v<cutlass::conv::KernelImplicitTmaWarpSpecializedSm90,
|
||||
typename CollectiveMainloop_::DispatchPolicy::Schedule>>>
|
||||
{
|
||||
@@ -90,6 +90,7 @@ public:
|
||||
using EpilogueArguments = typename CollectiveEpilogue::Arguments;
|
||||
using EpilogueParams = typename CollectiveEpilogue::Params;
|
||||
|
||||
using TileSchedulerTag = TileSchedulerTag_;
|
||||
static_assert(cute::is_void_v<TileSchedulerTag>,
|
||||
"TMA warp-specialized kernel does not support specializing the tile scheduler.");
|
||||
using TileScheduler = typename cutlass::gemm::kernel::detail::TileSchedulerSelector<
|
||||
@@ -144,7 +145,7 @@ public:
|
||||
to_underlying_arguments(Arguments const& args, void* workspace) {
|
||||
(void) workspace;
|
||||
auto mainloop_params = CollectiveMainloop::to_underlying_arguments(args.mainloop, workspace);
|
||||
auto problem_shape_MNKL = append<4>(mainloop_params.problem_shape, Int<1>{});
|
||||
auto problem_shape_MNKL = args.mainloop.problem_shape.get_transformed_problem_shape_MNKL();
|
||||
|
||||
return {
|
||||
mainloop_params,
|
||||
@@ -157,7 +158,7 @@ public:
|
||||
can_implement(Arguments const& args) {
|
||||
bool implementable = true;
|
||||
implementable &= CollectiveMainloop::can_implement(args.mainloop.problem_shape, args.mainloop);
|
||||
implementable &= CollectiveEpilogue::can_implement(args.mainloop.problem_shape.get_transformed_problem_shape_MNK(), args.epilogue);
|
||||
implementable &= CollectiveEpilogue::can_implement(args.mainloop.problem_shape.get_transformed_problem_shape_MNKL(), args.epilogue);
|
||||
return implementable;
|
||||
}
|
||||
|
||||
@@ -166,19 +167,17 @@ public:
|
||||
return 0;
|
||||
}
|
||||
|
||||
static cutlass::Status
|
||||
initialize_workspace(Arguments const& args, void* workspace = nullptr, cudaStream_t stream = nullptr,
|
||||
CudaHostAdapter* cuda_adapter = nullptr) {
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
// Computes the kernel launch grid shape based on runtime parameters
|
||||
static dim3
|
||||
get_grid_shape(Params const& params) {
|
||||
// The CONV mainloop params problem shape will be the cute::Shape<> rank-3 MNK tuple we want for grid planning
|
||||
// Although conv problems do not have an L mode, we add it here to comply with the scheduler API
|
||||
auto linear_problem_shape_MNKL = make_shape(
|
||||
size<0>(params.mainloop.problem_shape), // M mode is linearized.
|
||||
shape<1>(params.mainloop.problem_shape),
|
||||
shape<2>(params.mainloop.problem_shape),
|
||||
Int<1>{});
|
||||
|
||||
return cutlass::gemm::kernel::detail::PersistentTileSchedulerSm90::get_tiled_cta_shape_mnl(
|
||||
linear_problem_shape_MNKL, TileShape{}, ClusterShape{});
|
||||
params.mainloop.problem_shape, TileShape{}, ClusterShape{});
|
||||
}
|
||||
|
||||
static dim3
|
||||
@@ -205,14 +204,25 @@ public:
|
||||
Consumer = 1,
|
||||
};
|
||||
|
||||
enum class ProducerWarpRole {
|
||||
MainloopEpilogue = 0,
|
||||
Warp1 = 1,
|
||||
Warp2 = 2,
|
||||
Warp3 = 3
|
||||
};
|
||||
|
||||
// Kernel level shared memory storage
|
||||
SharedStorage& shared_storage = *reinterpret_cast<SharedStorage*>(smem_buf);
|
||||
|
||||
int thread_idx = int(threadIdx.x);
|
||||
int lane_idx = canonical_lane_idx();
|
||||
int warp_idx = canonical_warp_idx_sync();
|
||||
int warp_idx_in_warp_group = warp_idx % NumWarpsPerWarpGroup;
|
||||
int warp_group_thread_idx = thread_idx % NumThreadsPerWarpGroup;
|
||||
auto warp_group_role = WarpGroupRole(canonical_warp_group_idx());
|
||||
auto producer_warp_role = ProducerWarpRole(warp_idx_in_warp_group);
|
||||
int lane_predicate = cute::elect_one_sync();
|
||||
uint32_t block_rank_in_cluster = cute::block_rank_in_cluster();
|
||||
|
||||
// Issue Tma Descriptor Prefetch from a single thread
|
||||
if ((warp_idx == 0) && lane_predicate) {
|
||||
@@ -223,7 +233,7 @@ public:
|
||||
// Mainloop Load pipeline
|
||||
using MainloopPipeline = typename CollectiveMainloop::MainloopPipeline;
|
||||
typename MainloopPipeline::Params mainloop_pipeline_params;
|
||||
if (warp_group_role == WarpGroupRole::Producer) {
|
||||
if (warp_group_role == WarpGroupRole::Producer && producer_warp_role == ProducerWarpRole::MainloopEpilogue) {
|
||||
mainloop_pipeline_params.role = MainloopPipeline::ThreadCategory::Producer;
|
||||
}
|
||||
if (warp_group_role == WarpGroupRole::Consumer) {
|
||||
@@ -231,22 +241,24 @@ public:
|
||||
}
|
||||
mainloop_pipeline_params.is_leader = warp_group_thread_idx == 0;
|
||||
mainloop_pipeline_params.num_consumers = NumThreadsPerWarpGroup;
|
||||
mainloop_pipeline_params.transaction_bytes = CollectiveMainloop::TmaTransactionBytes;
|
||||
mainloop_pipeline_params.transaction_bytes = params.mainloop.tma_transaction_bytes;
|
||||
MainloopPipeline mainloop_pipeline(shared_storage.pipelines.mainloop, mainloop_pipeline_params, ClusterShape{});
|
||||
|
||||
// Epilogue Load pipeline
|
||||
using EpiLoadPipeline = typename CollectiveEpilogue::LoadPipeline;
|
||||
typename EpiLoadPipeline::Params epi_load_pipeline_params;
|
||||
if (warp_group_role == WarpGroupRole::Producer) {
|
||||
if (warp_group_role == WarpGroupRole::Producer && producer_warp_role == ProducerWarpRole::MainloopEpilogue) {
|
||||
epi_load_pipeline_params.role = EpiLoadPipeline::ThreadCategory::Producer;
|
||||
}
|
||||
if (warp_group_role == WarpGroupRole::Consumer) {
|
||||
epi_load_pipeline_params.role = EpiLoadPipeline::ThreadCategory::Consumer;
|
||||
}
|
||||
epi_load_pipeline_params.dst_blockid = cute::block_rank_in_cluster();
|
||||
epi_load_pipeline_params.producer_arv_count = 1; // 1 thread issues TMA load
|
||||
epi_load_pipeline_params.producer_arv_count = NumThreadsPerWarp;
|
||||
epi_load_pipeline_params.consumer_arv_count = NumThreadsPerWarpGroup;
|
||||
epi_load_pipeline_params.transaction_bytes = CollectiveEpilogue::TmaTransactionBytes;
|
||||
if constexpr (CollectiveEpilogue::RequiresTransactionBytes) {
|
||||
epi_load_pipeline_params.transaction_bytes = params.epilogue.tma_transaction_bytes;
|
||||
}
|
||||
EpiLoadPipeline epi_load_pipeline(shared_storage.pipelines.epi_load, epi_load_pipeline_params);
|
||||
|
||||
// Epilogue Store pipeline
|
||||
@@ -266,16 +278,26 @@ public:
|
||||
PipelineState epi_load_pipe_producer_state = cutlass::make_producer_start_state<EpiLoadPipeline>();
|
||||
PipelineState epi_store_pipe_producer_state = cutlass::make_producer_start_state<EpiStorePipeline>();
|
||||
|
||||
auto cluster_wait_fn = [&] () {
|
||||
// We need this to guarantee that the Pipeline init is visible
|
||||
// To all producers and consumer thread blocks in the Cluster
|
||||
if constexpr (size(ClusterShape{}) > 1) {
|
||||
cute::cluster_arrive_relaxed();
|
||||
return [] () { cute::cluster_wait(); };
|
||||
}
|
||||
else {
|
||||
__syncthreads();
|
||||
return [] () {}; // do nothing
|
||||
}
|
||||
} ();
|
||||
|
||||
// Separate out problem shape for convenience
|
||||
auto M = get<0>(params.mainloop.problem_shape);
|
||||
auto N = get<1>(params.mainloop.problem_shape);
|
||||
auto K = get<2>(params.mainloop.problem_shape);
|
||||
// output strides are coalesced so we linearize the output shape to match the shape/stride profiles
|
||||
auto linear_problem_shape_MNKL = make_shape(size(M), N, K, Int<1>{});
|
||||
auto problem_shape_MNKL = append<4>(params.mainloop.problem_shape, _1{});
|
||||
auto [M, N, K, L] = problem_shape_MNKL;
|
||||
|
||||
// TMA requires special handling of strides to deal with coord codomain mapping
|
||||
// Represent the full tensors -- get these from TMA
|
||||
Tensor mA_mk = params.mainloop.tma_load_a.get_tma_tensor(make_shape(M, size(K)));
|
||||
Tensor mA_mk = params.mainloop.tma_load_a.get_tma_tensor(make_shape(M, K));
|
||||
Tensor mB_nk = params.mainloop.tma_load_b.get_tma_tensor(make_shape(N, K));
|
||||
|
||||
// Get the appropriate blocks for this thread block -- potential for thread block locality
|
||||
@@ -288,7 +310,8 @@ public:
|
||||
|
||||
// Compute m_coord, n_coord, and l_coord with their post-tiled shapes
|
||||
auto m_coord = idx2crd(int(blockIdx.x), shape<2>(gA_mk));
|
||||
auto n_coord = idx2crd(int(blockIdx.y), shape<2>(gB_nk));
|
||||
auto n_coord = idx2crd(int(blockIdx.y), shape<2>(gB_nk), compact_col_major(shape<2>(gB_nk)));
|
||||
|
||||
// The output shape M is linearized so the output coord M here should also be linearized.
|
||||
auto output_tile_coord = make_coord(int(blockIdx.x), n_coord, _, Int<0>{});
|
||||
|
||||
@@ -300,51 +323,43 @@ public:
|
||||
auto k_tile_iter = cute::make_coord_iterator(shape<2>(gA));
|
||||
auto k_tile_count = size<2>(gA);
|
||||
|
||||
auto c_tile_count = CollectiveEpilogue::get_load_pipe_increment(cta_tile_shape);
|
||||
auto d_tile_count = CollectiveEpilogue::get_store_pipe_increment(cta_tile_shape);
|
||||
|
||||
// Make sure pipeline init is visible to all producers and consumer CTAs in cluster
|
||||
if constexpr (size(ClusterShape{}) > 1) {
|
||||
cute::cluster_arrive_relaxed();
|
||||
cute::cluster_wait();
|
||||
}
|
||||
else {
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
// In a warp specialized kernel, collectives expose data movement and compute operations separately
|
||||
CollectiveMainloop collective_mainloop;
|
||||
CollectiveEpilogue collective_epilogue{params.epilogue, shared_storage.tensors.epilogue};
|
||||
|
||||
if (warp_group_role == WarpGroupRole::Producer) {
|
||||
collective_mainloop.load(
|
||||
mainloop_pipeline,
|
||||
mainloop_pipe_producer_state,
|
||||
gA, params.mainloop.tma_load_a,
|
||||
gB, params.mainloop.tma_load_b,
|
||||
k_tile_iter, k_tile_count,
|
||||
thread_idx,
|
||||
shared_storage.tensors.mainloop
|
||||
);
|
||||
// Update starting mainloop pipeline state for the pipeline drain
|
||||
mainloop_pipe_producer_state.advance(k_tile_count);
|
||||
// Make sure mainloop consumer has been waited upon before issuing epilogue load
|
||||
collective_mainloop.load_tail(mainloop_pipeline, mainloop_pipe_producer_state);
|
||||
// Wait for all thread blocks in Cluster
|
||||
cluster_wait_fn();
|
||||
|
||||
if (collective_epilogue.is_producer_load_needed()) {
|
||||
collective_epilogue.load(
|
||||
epi_load_pipeline,
|
||||
epi_load_pipe_producer_state,
|
||||
linear_problem_shape_MNKL,
|
||||
cta_tile_shape,
|
||||
output_tile_coord,
|
||||
tiled_mma,
|
||||
warp_group_thread_idx,
|
||||
shared_storage.tensors.epilogue
|
||||
if (warp_group_role == WarpGroupRole::Producer) {
|
||||
if (producer_warp_role == ProducerWarpRole::MainloopEpilogue) {
|
||||
collective_mainloop.load(
|
||||
mainloop_pipeline,
|
||||
mainloop_pipe_producer_state,
|
||||
gA, params.mainloop.tma_load_a,
|
||||
gB, params.mainloop.tma_load_b,
|
||||
k_tile_iter, k_tile_count,
|
||||
lane_idx,
|
||||
block_rank_in_cluster,
|
||||
shared_storage.tensors.mainloop
|
||||
);
|
||||
// Update starting load pipeline state for the pipeline drain
|
||||
epi_load_pipe_producer_state.advance(c_tile_count);
|
||||
collective_epilogue.load_tail(epi_load_pipeline, epi_load_pipe_producer_state);
|
||||
// Update starting mainloop pipeline state for the pipeline drain
|
||||
mainloop_pipe_producer_state.advance(k_tile_count);
|
||||
// Make sure mainloop consumer has been waited upon before issuing epilogue load
|
||||
collective_mainloop.load_tail(mainloop_pipeline, mainloop_pipe_producer_state);
|
||||
|
||||
if (collective_epilogue.is_producer_load_needed()) {
|
||||
epi_load_pipe_producer_state = collective_epilogue.load(
|
||||
epi_load_pipeline,
|
||||
epi_load_pipe_producer_state,
|
||||
problem_shape_MNKL,
|
||||
cta_tile_shape,
|
||||
output_tile_coord,
|
||||
tiled_mma,
|
||||
lane_idx,
|
||||
shared_storage.tensors.epilogue
|
||||
);
|
||||
collective_epilogue.load_tail(epi_load_pipeline, epi_load_pipe_producer_state);
|
||||
}
|
||||
}
|
||||
}
|
||||
else if (warp_group_role == WarpGroupRole::Consumer) {
|
||||
@@ -368,12 +383,13 @@ public:
|
||||
);
|
||||
|
||||
// Epilogue and write to gD
|
||||
auto [epi_load_pipe_consumer_state_next, epi_store_pipe_producer_state_next] =
|
||||
collective_epilogue.store(
|
||||
epi_load_pipeline,
|
||||
epi_load_pipe_consumer_state,
|
||||
epi_store_pipeline,
|
||||
epi_store_pipe_producer_state,
|
||||
linear_problem_shape_MNKL,
|
||||
problem_shape_MNKL,
|
||||
cta_tile_shape,
|
||||
output_tile_coord,
|
||||
accumulators,
|
||||
@@ -381,6 +397,13 @@ public:
|
||||
warp_group_thread_idx,
|
||||
shared_storage.tensors.epilogue
|
||||
);
|
||||
|
||||
collective_epilogue.store_tail(
|
||||
epi_load_pipeline,
|
||||
epi_load_pipe_consumer_state_next,
|
||||
epi_store_pipeline,
|
||||
epi_store_pipe_producer_state_next
|
||||
);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
+1
-1
@@ -251,7 +251,7 @@ public:
|
||||
static Status can_implement(Conv3dProblemSize const &problem_size) {
|
||||
|
||||
// check alignment constraint on iterator's contiguous dimension
|
||||
if (problem_size.C % (128/sizeof_bits<Element>::value)) {
|
||||
if (problem_size.C % AccessType::kElements) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
|
||||
+1
-1
@@ -272,7 +272,7 @@ public:
|
||||
static Status can_implement(Conv3dProblemSize const &problem_size) {
|
||||
|
||||
// check alignment constraint on iterator's contiguous dimension
|
||||
if (problem_size.C % (128/sizeof_bits<Element>::value)) {
|
||||
if (problem_size.C % AccessType::kElements) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
|
||||
+1
-1
@@ -325,7 +325,7 @@ public:
|
||||
static Status can_implement(ConvProblemSize const &problem_size) {
|
||||
|
||||
// check alignment constraint on iterator's contiguous dimension
|
||||
if (problem_size.K % (128/sizeof_bits<Element>::value)) {
|
||||
if (problem_size.K % AccessType::kElements) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
|
||||
+1
-1
@@ -466,7 +466,7 @@ public:
|
||||
}
|
||||
|
||||
// check alignment constraint on iterator's contiguous dimension
|
||||
if (problem_size.K % (128/sizeof_bits<Element>::value)) {
|
||||
if (problem_size.K % AccessType::kElements) {
|
||||
return Status::kErrorNotSupported;
|
||||
}
|
||||
|
||||
|
||||
+1
-1
@@ -272,7 +272,7 @@ public:
|
||||
static Status can_implement(ConvProblemSize const &problem_size) {
|
||||
|
||||
// check alignment constraint on iterator's contiguous dimension
|
||||
if (problem_size.C % (128/sizeof_bits<Element>::value)) {
|
||||
if (problem_size.C % AccessType::kElements) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
|
||||
+1
-1
@@ -455,7 +455,7 @@ public:
|
||||
static Status can_implement(Conv3dProblemSize const &problem_size) {
|
||||
|
||||
// check alignment constraint on iterator's contiguous dimension
|
||||
if (problem_size.C % (128/sizeof_bits<Element>::value)) {
|
||||
if (problem_size.C % AccessType::kElements) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
|
||||
+1
-2
@@ -238,11 +238,10 @@ public:
|
||||
/// Determines whether the Implicit GEMM can execute the given problem.
|
||||
CUTLASS_HOST_DEVICE
|
||||
static Status can_implement(ConvProblemSize const &problem_size) {
|
||||
|
||||
auto input_channels = (IsDeconv ? problem_size.K : problem_size.C);
|
||||
auto output_channels = (IsDeconv ? problem_size.C : problem_size.K);
|
||||
// check alignment constraint on iterator's contiguous dimension
|
||||
if (input_channels % (128/sizeof_bits<Element>::value)) {
|
||||
if (input_channels % AccessType::kElements) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
return Status::kSuccess;
|
||||
|
||||
+1
-3
@@ -260,14 +260,12 @@ public:
|
||||
/// Determines whether the Implicit GEMM can execute the given problem.
|
||||
CUTLASS_HOST_DEVICE
|
||||
static Status can_implement(Conv3dProblemSize const &problem_size) {
|
||||
|
||||
auto input_channels = (IsDeconv ? problem_size.K : problem_size.C);
|
||||
|
||||
// check alignment constraint on iterator's contiguous dimension
|
||||
if (input_channels % (128/sizeof_bits<Element>::value)) {
|
||||
if (input_channels % AccessType::kElements) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
};
|
||||
|
||||
+1
-1
@@ -270,7 +270,7 @@ public:
|
||||
static Status can_implement(Conv3dProblemSize const &problem_size) {
|
||||
|
||||
// check alignment constraint on iterator's contiguous dimension
|
||||
if (problem_size.C % (128/sizeof_bits<Element>::value)) {
|
||||
if (problem_size.C % AccessType::kElements) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
|
||||
+2
-2
@@ -250,7 +250,7 @@ public:
|
||||
fast_divmod(p, q, residual, problem_size_.Q, params_.q_mul, params_.q_shr);
|
||||
|
||||
int d = z * problem_size_.stride_d + precomputed_filter_t_[iteration_contiguous_];
|
||||
int h = p * problem_size_.stride_h + precomputed_filter_r_[iteration_contiguous_];;
|
||||
int h = p * problem_size_.stride_h + precomputed_filter_r_[iteration_contiguous_];
|
||||
int w = q * problem_size_.stride_w + precomputed_filter_s_[iteration_contiguous_];
|
||||
|
||||
return TensorCoord(n, d, h, w, filter_c_[iteration_contiguous_]);
|
||||
@@ -300,7 +300,7 @@ public:
|
||||
static Status can_implement(Conv3dProblemSize const &problem_size) {
|
||||
|
||||
// check alignment constraint on iterator's contiguous dimension
|
||||
if (problem_size.C % (128/sizeof_bits<Element>::value)) {
|
||||
if (problem_size.C % AccessType::kElements) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
|
||||
+1
-1
@@ -248,7 +248,7 @@ public:
|
||||
static Status can_implement(Conv3dProblemSize const &problem_size) {
|
||||
|
||||
// check alignment constraint on iterator's contiguous dimension
|
||||
if (problem_size.K % (128/sizeof_bits<Element>::value)) {
|
||||
if (problem_size.K % AccessType::kElements) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
|
||||
+1
-1
@@ -291,7 +291,7 @@ public:
|
||||
static Status can_implement(Conv3dProblemSize const &problem_size) {
|
||||
|
||||
// check alignment constraint on iterator's contiguous dimension
|
||||
if (problem_size.K % (128/sizeof_bits<Element>::value)) {
|
||||
if (problem_size.K % AccessType::kElements) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user