fix allgather op select

This commit is contained in:
yuchaojie 2021-06-25 17:25:09 +08:00
parent 9886aa613e
commit 9d9c014ed1
3 changed files with 42 additions and 30 deletions

View File

@ -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()) {

View File

@ -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

View File

@ -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);