Browse Source

Transdata

pull/1211/head
zk 5 years ago
parent
commit
3925716297
3 changed files with 30 additions and 3 deletions
  1. +17
    -1
      ge/common/formats/format_transfers/format_transfer_fractal_z.cc
  2. +1
    -0
      ge/common/formats/format_transfers/format_transfer_fractal_z.h
  3. +12
    -2
      ge/host_kernels/transdata_kernel.cc

+ 17
- 1
ge/common/formats/format_transfers/format_transfer_fractal_z.cc View File

@@ -317,7 +317,7 @@ Status TransFormatHwcnToFzWithGroups(const TransArgs &args, TransResult &result)
h * w_dim * c_dim * n_dim + w * c_dim * n_dim +
c * n_dim + src_co;
char *dst_data = reinterpret_cast<char *>(dst.get() + dst_inx * data_size);
const char *src_data = reinterpret_cast<const char *>(args.data + src_idx * data_size);
const char *src_data = reinterpret_cast<const char *>(args.data + srx_inx * data_size);
for (int64_t index = 0; index < data_size; index++) {
*dst_data++ = *src_data++;
}
@@ -481,6 +481,22 @@ Status TransFormatNhwcToFz(const TransArgs &args, TransResult &result) {
}
} // namespace

Status FormatTransferFractalZ::TransFormat(const TransArgs &args, TransResult &result,
int64_t groups) {
GELOGD("Begin to trans format from %s to %s, src shape %s, data type %s, dst shape %s",
TypeUtils::FormatToSerialString(args.src_format).c_str(),
TypeUtils::FormatToSerialString(args.dst_format).c_str(), ShapeToString(args.src_shape).c_str(),
TypeUtils::DataTypeToSerialString(args.src_data_type).c_str(), ShapeToString(args.dst_shape).c_str());
std::vector<int64_t> expect_shape;
auto ret = TransShape(args.src_format, args.src_shape, args.src_data_type, args.dst_format, expect_shape);
if (ret != SUCCESS) {
return ret;
}
if (!IsTransShapeDstCorrect(args, expect_shape)) {
return PARAM_INVALID;
}
return TransFormatHwcnToFzWithGroups(args, result);
}
Status FormatTransferFractalZ::TransFormat(const TransArgs &args, TransResult &result) {
GELOGD("Begin to trans format from %s to %s, src shape %s, data type %s, dst shape %s",
TypeUtils::FormatToSerialString(args.src_format).c_str(),


+ 1
- 0
ge/common/formats/format_transfers/format_transfer_fractal_z.h View File

@@ -26,6 +26,7 @@ namespace formats {
class FormatTransferFractalZ : public FormatTransfer {
public:
Status TransFormat(const TransArgs &args, TransResult &result) override;
Status TransFormat(const TransArgs &args, TransResult &result, int64_t groups) override;
Status TransShape(Format src_format, const std::vector<int64_t> &src_shape, DataType data_type, Format dst_format,
std::vector<int64_t> &dst_shape) override;
};


+ 12
- 2
ge/host_kernels/transdata_kernel.cc View File

@@ -82,16 +82,17 @@ Status TransdataKernel::Compute(const OpDescPtr op_desc_ptr, const std::vector<C
const auto &data_shape = op_desc->GetShape().GetDims();
const auto &data_format = op_desc->GetFormat();
const auto &data_type = op_desc->GetDataType();
const in64_t groups = op_desc_ptr->GetAttr("groups", groups)
GELOGD(
"current node %s, format %s, input shape %s, data type %s, weight format %s, shape %s, data type %s. "
"output format %s, shape %s, data type %s",
"output format %s, shape %s, data type %s, groups %d",
op_desc_ptr->GetName().c_str(), TypeUtils::FormatToSerialString(src_format).c_str(),
formats::ShapeToString(src_shape).c_str(), TypeUtils::DataTypeToSerialString(src_data_type).c_str(),
TypeUtils::FormatToSerialString(const_weight_ptr->GetTensorDesc().GetFormat()).c_str(),
formats::ShapeToString(const_weight_ptr->GetTensorDesc().GetShape()).c_str(),
TypeUtils::DataTypeToSerialString(const_weight_ptr->GetTensorDesc().GetDataType()).c_str(),
TypeUtils::FormatToSerialString(data_format).c_str(), formats::ShapeToString(data_shape).c_str(),
TypeUtils::DataTypeToSerialString(data_type).c_str());
TypeUtils::DataTypeToSerialString(data_type).c_str(), groups);

const uint8_t *src_data = const_weight_ptr->GetData().data();
const formats::TransArgs trans_args{src_data, src_format, data_format, src_shape, data_shape, src_data_type};
@@ -113,6 +114,15 @@ Status TransdataKernel::Compute(const OpDescPtr op_desc_ptr, const std::vector<C
GELOGI("CheckSize failed, input size is not equal to weight size");
return NOT_CHANGED;
}
if((src_format == FOMAT_HWCN) && (data_format == FORMAT_FRACTAL_Z_3D)) {
if (formats::TransFormat(trans_args, trans_result , groups) != SUCCESS) {
GELOGW("Failed to trans formats from %s to %s, shape %s to %s, data type %s",
TypeUtils::FormatToSerialString(src_format).c_str(), TypeUtils::FormatToSerialString(data_format).c_str(),
formats::ShapeToString(src_shape).c_str(), formats::ShapeToString(data_shape).c_str(),
TypeUtils::DataTypeToSerialString(src_data_type).c_str());
return NOT_CHANGED;
}
}
if (formats::TransFormat(trans_args, trans_result) != SUCCESS) {
GELOGW("Failed to trans formats from %s to %s, shape %s to %s, data type %s",
TypeUtils::FormatToSerialString(src_format).c_str(), TypeUtils::FormatToSerialString(data_format).c_str(),


Loading…
Cancel
Save