forked from huawei/mindspore2022
fix allgather op select
This commit is contained in:
parent
9886aa613e
commit
9d9c014ed1
|
|
@ -18,6 +18,7 @@
|
|||
#include <memory>
|
||||
#include <algorithm>
|
||||
#include <set>
|
||||
#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<std::string> 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()) {
|
||||
|
|
|
|||
|
|
@ -380,32 +380,6 @@ std::vector<size_t> ChannelLastDeviceShape(const std::vector<size_t> &shape) {
|
|||
return device_shape;
|
||||
}
|
||||
|
||||
std::vector<size_t> PaddingShapeTo4dByDefault(const std::vector<size_t> &shape) {
|
||||
std::vector<size_t> 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<size_t> FracZDeviceShapeWithGroups(const std::vector<size_t> &shape, const int64_t groups = 1) {
|
||||
if (!CheckDims(shape)) {
|
||||
MS_LOG(EXCEPTION) << "Check dims failed.";
|
||||
|
|
@ -579,7 +553,7 @@ std::vector<size_t> PaddingShapeTo4d(const std::vector<size_t> &shape, const std
|
|||
std::vector<Axis> padding_axis;
|
||||
StringToAxisVector4D(padding_str, &padding_axis);
|
||||
if (padding_axis.empty() || shape.size() != padding_axis.size()) {
|
||||
return PaddingShapeTo4dByDefault(shape);
|
||||
return PaddingShapeTo4dDefault(shape);
|
||||
}
|
||||
std::vector<size_t> shape_4d(kNchwDims, 1);
|
||||
for (size_t index = 0; index < padding_axis.size(); index++) {
|
||||
|
|
@ -620,6 +594,32 @@ std::vector<size_t> PaddingShapeTo5dDefault(const std::vector<size_t> &shape) {
|
|||
return shape_5d;
|
||||
}
|
||||
|
||||
std::vector<size_t> PaddingShapeTo4dDefault(const std::vector<size_t> &shape) {
|
||||
std::vector<size_t> 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<size_t> TransShapeToDevice(const std::vector<size_t> &shape, const std::string &format,
|
||||
const int64_t groups) {
|
||||
using DeviceShapeTransfer = std::function<std::vector<size_t>(const std::vector<size_t> &)>;
|
||||
|
|
@ -675,7 +675,7 @@ std::vector<size_t> TransShapeToDevice(const std::vector<size_t> &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
|
||||
|
|
|
|||
|
|
@ -60,6 +60,7 @@ std::vector<size_t> PaddingShape(const std::vector<size_t> &shape, const std::st
|
|||
std::vector<size_t> PaddingShapeTo4d(const std::vector<size_t> &shape, const std::string &padding_axis = {""});
|
||||
std::vector<size_t> PaddingShapeTo5d(const std::vector<size_t> &shape, const std::string &padding_axis = {""});
|
||||
std::vector<size_t> PaddingShapeTo5dDefault(const std::vector<size_t> &shape);
|
||||
std::vector<size_t> PaddingShapeTo4dDefault(const std::vector<size_t> &shape);
|
||||
void StringToAxisVector4D(const std::string &reshape_type_str, std::vector<Axis> *reshape_type_vec);
|
||||
void StringToAxisVector5D(const std::string &reshape_type_str, std::vector<Axis5D> *reshape_type_vec);
|
||||
ShapeVector GetRuntimePaddingShape(const AnfNodePtr &node, size_t index);
|
||||
|
|
|
|||
Loading…
Reference in New Issue