提交 488e889f 编写于 作者: P phlrain

fix split infer shape; test=develop

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