update 3.8 v2 (#2112)
* update 3.8 v2 * update 3.8 --------- Co-authored-by: yuzhai <yuzhai@nvidia.com>
This commit is contained in:
@@ -106,7 +106,6 @@ struct CollectiveConv<
|
||||
|
||||
using ProblemShape = ConvProblemShape<ConvOp, NumSpatialDimensions>;
|
||||
|
||||
// TODO: move pipeline mode tiling into the collective setup phase instead
|
||||
static_assert(rank(SmemLayoutA{}) == 3, "SmemLayout must be rank 3 (M/N, K, PIPE)");
|
||||
static_assert((size<0>(TileShape{}) == size<0>(SmemLayoutA{})), "SmemLayout must be compatible with the tile shape.");
|
||||
static_assert((size<2>(TileShape{}) == size<1>(SmemLayoutA{})), "SmemLayout must be compatible with the tile shape.");
|
||||
|
||||
@@ -255,23 +255,27 @@ public:
|
||||
CUTLASS_HOST_DEVICE
|
||||
int64_t activation_size() const {
|
||||
|
||||
return (N * H * W * C);
|
||||
return static_cast<int64_t>(N) * static_cast<int64_t>(H) *
|
||||
static_cast<int64_t>(W) * static_cast<int64_t>(C);
|
||||
}
|
||||
|
||||
/// Returns filter size in number of elements
|
||||
CUTLASS_HOST_DEVICE
|
||||
int64_t filter_size() const {
|
||||
|
||||
return (K * R * S * C / groups);
|
||||
return static_cast<int64_t>(K) * static_cast<int64_t>(R) *
|
||||
static_cast<int64_t>(S) * static_cast<int64_t>(C) /
|
||||
static_cast<int64_t>(groups);
|
||||
}
|
||||
|
||||
/// Returns output size in number of elements
|
||||
CUTLASS_HOST_DEVICE
|
||||
int64_t output_size() const {
|
||||
|
||||
return (N * P * Q * K);
|
||||
return static_cast<int64_t>(N) * static_cast<int64_t>(P) *
|
||||
static_cast<int64_t>(Q) * static_cast<int64_t>(K);
|
||||
}
|
||||
|
||||
|
||||
/// Returns padding as Tensor4DCoord
|
||||
CUTLASS_HOST_DEVICE
|
||||
cutlass::Tensor4DCoord padding() const {
|
||||
|
||||
@@ -285,21 +285,27 @@ public:
|
||||
CUTLASS_HOST_DEVICE
|
||||
int64_t activation_size() const {
|
||||
|
||||
return (N * D * H * W * C);
|
||||
return static_cast<int64_t>(N) * static_cast<int64_t>(D) *
|
||||
static_cast<int64_t>(H) * static_cast<int64_t>(W) *
|
||||
static_cast<int64_t>(C);
|
||||
}
|
||||
|
||||
/// Returns filter size in number of elements
|
||||
CUTLASS_HOST_DEVICE
|
||||
int64_t filter_size() const {
|
||||
|
||||
return (K * T * R * S * C);
|
||||
return static_cast<int64_t>(K) * static_cast<int64_t>(T) *
|
||||
static_cast<int64_t>(R) * static_cast<int64_t>(S) *
|
||||
static_cast<int64_t>(C);
|
||||
}
|
||||
|
||||
/// Returns output size in number of elements
|
||||
CUTLASS_HOST_DEVICE
|
||||
int64_t output_size() const {
|
||||
|
||||
return (N * Z * P * Q * K);
|
||||
return static_cast<int64_t>(N) * static_cast<int64_t>(Z) *
|
||||
static_cast<int64_t>(P) * static_cast<int64_t>(Q) *
|
||||
static_cast<int64_t>(K);
|
||||
}
|
||||
|
||||
/// Returns padding as Coord3D
|
||||
|
||||
@@ -114,6 +114,33 @@ public:
|
||||
return status;
|
||||
}
|
||||
|
||||
// Check that tensor sizes don't exceed maximum supported size
|
||||
if (kConvolutionalOperator == conv::Operator::kFprop) {
|
||||
if (args.problem_size.activation_size() * sizeof(ElementA) >=
|
||||
(1ull << 31) ||
|
||||
args.problem_size.filter_size() * sizeof(ElementB) >= (1ull << 31) ||
|
||||
args.problem_size.output_size() * sizeof(ElementC) >= (1ull << 31)) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
}
|
||||
else if (kConvolutionalOperator == conv::Operator::kDgrad ||
|
||||
kConvolutionalOperator == conv::Operator::kDeconv) {
|
||||
if (args.problem_size.activation_size() * sizeof(ElementC) >=
|
||||
(1ull << 31) ||
|
||||
args.problem_size.filter_size() * sizeof(ElementB) >= (1ull << 31) ||
|
||||
args.problem_size.output_size() * sizeof(ElementA) >= (1ull << 31)) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
}
|
||||
else if (kConvolutionalOperator == conv::Operator::kWgrad) {
|
||||
if (args.problem_size.activation_size() * sizeof(ElementB) >=
|
||||
(1ull << 31) ||
|
||||
args.problem_size.filter_size() * sizeof(ElementC) >= (1ull << 31) ||
|
||||
args.problem_size.output_size() * sizeof(ElementA) >= (1ull << 31)) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
}
|
||||
|
||||
// check group conv constraint
|
||||
if (args.problem_size.groups != 1) {
|
||||
if (kGroupMode == conv::GroupMode::kNone) {
|
||||
|
||||
Reference in New Issue
Block a user