Skip to content

Commit

Permalink
fix cond create undefined var in global block (PaddlePaddle#58065)
Browse files Browse the repository at this point in the history
  • Loading branch information
2742195759 authored Oct 17, 2023
1 parent ab0886d commit 30830d9
Show file tree
Hide file tree
Showing 2 changed files with 33 additions and 0 deletions.
15 changes: 15 additions & 0 deletions python/paddle/jit/dy2static/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -178,6 +178,21 @@ def data_layer_not_check(name, shape, dtype='float32', lod_level=0):
)


def create_undefined_variable_local():
helper = LayerHelper('create_undefined_variable', **locals())
var = helper.create_variable(
name=unique_name.generate("undefined_var"),
shape=[1],
dtype="float64",
type=core.VarDesc.VarType.LOD_TENSOR,
stop_gradient=False,
is_data=True,
need_check_feed=False,
)
paddle.assign(RETURN_NO_VALUE_MAGIC_NUM, var)
return var


def create_undefined_variable():
var = data_layer_not_check(
unique_name.generate("undefined_var"), [1], "float64"
Expand Down
18 changes: 18 additions & 0 deletions python/paddle/static/nn/control_flow.py
Original file line number Diff line number Diff line change
Expand Up @@ -1224,6 +1224,9 @@ def cond(pred, true_fn=None, false_fn=None, name=None, return_names=None):
with true_cond_block.block():
origin_true_output = true_fn()
if origin_true_output is not None:
origin_true_output = map_structure(
create_undefined_var_in_subblock, origin_true_output
)
true_output = map_structure(
copy_to_parent_func, origin_true_output
)
Expand All @@ -1240,6 +1243,9 @@ def cond(pred, true_fn=None, false_fn=None, name=None, return_names=None):
with false_cond_block.block():
origin_false_output = false_fn()
if origin_false_output is not None:
origin_false_output = map_structure(
create_undefined_var_in_subblock, origin_false_output
)
false_output = map_structure(
copy_to_parent_func, origin_false_output
)
Expand Down Expand Up @@ -1356,6 +1362,18 @@ def merge_every_var_list(false_vars, true_vars, name):
return merged_output


def create_undefined_var_in_subblock(var):
# to make sure the undefined var created in subblock.
from paddle.jit.dy2static.utils import (
UndefinedVar,
create_undefined_variable_local,
)

if isinstance(var, UndefinedVar):
var = create_undefined_variable_local()
return var


def copy_var_to_parent_block(var, layer_helper):
if not isinstance(var, Variable):
return var
Expand Down

0 comments on commit 30830d9

Please sign in to comment.