未验证 提交 c96ee47d 编写于 作者: H Hongyu Liu 提交者: GitHub

Merge pull request #16797 from phlrain/fix_split

Fix split
...@@ -39,6 +39,7 @@ class SplitOp : public framework::OperatorWithKernel { ...@@ -39,6 +39,7 @@ class SplitOp : public framework::OperatorWithKernel {
if (num > 0) { if (num > 0) {
int64_t in_axis_dim = in_dims[axis]; int64_t in_axis_dim = in_dims[axis];
if (ctx->IsRuntime() || in_axis_dim > 0) {
PADDLE_ENFORCE_EQ(in_axis_dim % num, 0, PADDLE_ENFORCE_EQ(in_axis_dim % num, 0,
"tensor split does not result" "tensor split does not result"
" in an equal division"); " in an equal division");
...@@ -48,6 +49,13 @@ class SplitOp : public framework::OperatorWithKernel { ...@@ -48,6 +49,13 @@ class SplitOp : public framework::OperatorWithKernel {
dim[axis] = out_axis_dim; dim[axis] = out_axis_dim;
outs_dims.push_back(dim); outs_dims.push_back(dim);
} }
} else {
for (size_t i = 0; i < outs_number; ++i) {
auto dim = in_dims;
dim[axis] = -1;
outs_dims.push_back(dim);
}
}
} else if (sections.size() > 0) { } else if (sections.size() > 0) {
PADDLE_ENFORCE_EQ(sections.size(), outs_number, PADDLE_ENFORCE_EQ(sections.size(), outs_number,
"tensor split sections size" "tensor split sections size"
......
Markdown is supported
0% .
You are about to add 0 people to the discussion. Proceed with caution.
先完成此消息的编辑!
想要评论请 注册