diff --git a/mindspore/ccsrc/backend/kernel_compiler/hccl/hccl_kernel_metadata.cc b/mindspore/ccsrc/backend/kernel_compiler/hccl/hccl_kernel_metadata.cc index 75e1f3b9853..92d26431f8e 100755 --- a/mindspore/ccsrc/backend/kernel_compiler/hccl/hccl_kernel_metadata.cc +++ b/mindspore/ccsrc/backend/kernel_compiler/hccl/hccl_kernel_metadata.cc @@ -18,6 +18,7 @@ #include #include #include +#include "common/trans.h" #include "utils/utils.h" #include "backend/kernel_compiler/hccl/hcom_util.h" #include "backend/session/anf_runtime_algorithm.h" @@ -26,6 +27,8 @@ namespace mindspore { namespace kernel { namespace { +constexpr size_t N_nchw = 0; +constexpr size_t C_nchw = 1; std::string GetKernelFormat(const CNodePtr &kernel_node, size_t index) { const std::set kReduceNoSupportedSet = {kOpFormat_FRAC_Z, kOpFormat_FRACTAL_Z_C04, kOpFormat_C1HWNCoC0}; auto op_name = AnfAlgo::GetCNodeName(kernel_node); @@ -41,7 +44,14 @@ std::string GetKernelFormat(const CNodePtr &kernel_node, size_t index) { if (op_name != kReduceScatter && op_name != kAllGatherOpName) { return format; } - if (format == kOpFormat_FRAC_NZ && AnfAlgo::GetPrevNodeOutputInferShape(kernel_node, index).size() <= 2) { + auto input_shape = AnfAlgo::GetPrevNodeOutputInferShape(kernel_node, index); + if (op_name == kAllGatherOpName) { + auto pad_shape = trans::PaddingShapeTo4dDefault(input_shape); + if (pad_shape[N_nchw] % kCubeSize != 0 || pad_shape[C_nchw] % kCubeSize != 0) { + return kOpFormat_DEFAULT; + } + } + if (format == kOpFormat_FRAC_NZ && input_shape.size() <= 2) { return kOpFormat_DEFAULT; } if (kReduceNoSupportedSet.find(format) != kReduceNoSupportedSet.end()) { diff --git a/mindspore/ccsrc/common/trans.cc b/mindspore/ccsrc/common/trans.cc index e6732e7210b..bb4dd2a9f92 100644 --- a/mindspore/ccsrc/common/trans.cc +++ b/mindspore/ccsrc/common/trans.cc @@ -380,32 +380,6 @@ std::vector ChannelLastDeviceShape(const std::vector &shape) { return device_shape; } -std::vector PaddingShapeTo4dByDefault(const std::vector &shape) { - std::vector shape_4d(kNchwDims, 1); - switch (shape.size()) { - case 0: - return shape_4d; - case 1: - shape_4d[kC] = shape[kN]; - break; - case 2: - shape_4d[kC] = shape[kN]; - shape_4d[kH] = shape[kC]; - break; - case 3: - shape_4d[kC] = shape[kN]; - shape_4d[kH] = shape[kC]; - shape_4d[kW] = shape[kH]; - break; - case 4: - std::copy(shape.begin(), shape.end(), shape_4d.begin()); - break; - default: - MS_LOG(EXCEPTION) << "Unexpected shape size = " << shape.size(); - } - return shape_4d; -} - std::vector FracZDeviceShapeWithGroups(const std::vector &shape, const int64_t groups = 1) { if (!CheckDims(shape)) { MS_LOG(EXCEPTION) << "Check dims failed."; @@ -579,7 +553,7 @@ std::vector PaddingShapeTo4d(const std::vector &shape, const std std::vector padding_axis; StringToAxisVector4D(padding_str, &padding_axis); if (padding_axis.empty() || shape.size() != padding_axis.size()) { - return PaddingShapeTo4dByDefault(shape); + return PaddingShapeTo4dDefault(shape); } std::vector shape_4d(kNchwDims, 1); for (size_t index = 0; index < padding_axis.size(); index++) { @@ -620,6 +594,32 @@ std::vector PaddingShapeTo5dDefault(const std::vector &shape) { return shape_5d; } +std::vector PaddingShapeTo4dDefault(const std::vector &shape) { + std::vector shape_4d(kNchwDims, 1); + switch (shape.size()) { + case 0: + return shape_4d; + case 1: + shape_4d[kC] = shape[kN]; + break; + case 2: + shape_4d[kC] = shape[kN]; + shape_4d[kH] = shape[kC]; + break; + case 3: + shape_4d[kC] = shape[kN]; + shape_4d[kH] = shape[kC]; + shape_4d[kW] = shape[kH]; + break; + case 4: + std::copy(shape.begin(), shape.end(), shape_4d.begin()); + break; + default: + MS_LOG(EXCEPTION) << "Unexpected shape size = " << shape.size(); + } + return shape_4d; +} + std::vector TransShapeToDevice(const std::vector &shape, const std::string &format, const int64_t groups) { using DeviceShapeTransfer = std::function(const std::vector &)>; @@ -675,7 +675,7 @@ std::vector TransShapeToDevice(const std::vector &shape, const s } if (format != kOpFormat_ChannelLast && shape.size() != kNchwDims && k3DFormatSet.find(format) == k3DFormatSet.end()) { MS_LOG(WARNING) << "Get Device Shape using a shape size is less than 4 ,should be Padding shape by Default firstly"; - temp_shape = PaddingShapeTo4dByDefault(shape); + temp_shape = PaddingShapeTo4dDefault(shape); } if (shape.size() != kNcdhw && k3DFormatSet.find(format) != k3DFormatSet.end()) { temp_shape = PaddingShapeTo5dDefault(shape); @@ -1587,7 +1587,6 @@ bool FracZ3DToNcdhw(const FormatArgs &args, void *result) { } bool NchwFracZTransWithGroups(const FormatArgs &args, void *result, bool to_device, int64_t groups) { - MS_LOG(DEBUG) << "Trans format from nchw to frac_z"; MS_EXCEPTION_IF_NULL(result); if (args.host_shape.size() != kNchwDims) { MS_LOG(ERROR) << "Invalid host shape, host shape dims:" << args.host_shape.size() << ", expect dims:" << kNchwDims; @@ -1656,10 +1655,12 @@ bool NchwFracZTransWithGroups(const FormatArgs &args, void *result, bool to_devi } bool NchwToFracZWithGroups(const FormatArgs &args, void *result, int64_t groups) { + MS_LOG(DEBUG) << "Trans format from nchw to frac_z with groups=" << groups; return NchwFracZTransWithGroups(args, result, true, groups); } bool FracZToNchwWithGroups(const FormatArgs &args, void *result, int64_t groups) { + MS_LOG(DEBUG) << "Trans format from frac_z to nchw with groups=" << groups; return NchwFracZTransWithGroups(args, result, false, groups); } } // namespace trans diff --git a/mindspore/ccsrc/common/trans.h b/mindspore/ccsrc/common/trans.h index b8992c07383..0f07a742282 100644 --- a/mindspore/ccsrc/common/trans.h +++ b/mindspore/ccsrc/common/trans.h @@ -60,6 +60,7 @@ std::vector PaddingShape(const std::vector &shape, const std::st std::vector PaddingShapeTo4d(const std::vector &shape, const std::string &padding_axis = {""}); std::vector PaddingShapeTo5d(const std::vector &shape, const std::string &padding_axis = {""}); std::vector PaddingShapeTo5dDefault(const std::vector &shape); +std::vector PaddingShapeTo4dDefault(const std::vector &shape); void StringToAxisVector4D(const std::string &reshape_type_str, std::vector *reshape_type_vec); void StringToAxisVector5D(const std::string &reshape_type_str, std::vector *reshape_type_vec); ShapeVector GetRuntimePaddingShape(const AnfNodePtr &node, size_t index);