PyTorch1.7.1的torch.Assert语句如何替换为PyTorch1.5.0支持的等价语句?
PyTorch1.5.0 替代 torch.Assert 的实现方案
torch.Assert 是 PyTorch1.7+ 新增的面向 TorchScript 符号图的断言接口,在 1.5.0 版本中可以通过以下两种方式实现等价功能:
方案1:Python 原生 assert 语句(非 TorchScript 导出场景首选)
如果你的代码不需要导出为 TorchScript 格式,直接用 Python 原生 assert 即可,运行时语义和 torch.Assert 完全一致:
# 原 1.7+ 代码 # torch.Assert(condition, message="error info") # 1.5.0 替换写法 assert condition, "error info"
注意:原生 assert 仅会在 Python 运行时生效,如果需要做 TorchScript 静态图导出时的断言校验,需要用第二种方案
方案2:自定义 TorchScript 兼容的断言函数(需要导出 TorchScript 场景使用)
如果你的代码需要导出 TorchScript 静态图,可以基于 1.5.0 版本已经存在的内部接口 torch._assert 封装等价功能:
def torch_assert(condition, message=""): # 等价于 torch.Assert 的符号断言能力,兼容 PyTorch1.5.0 if torch.jit.is_scripting(): torch._assert(condition, message) else: assert condition, message
调用时直接替换原 torch.Assert 调用即可:
# 原写法 # torch.Assert(x > 0, message="x must be positive") # 替换后写法 torch_assert(x > 0, message="x must be positive")
内容的提问来源于stack exchange,提问作者Zeke
相关产品推荐
相关产品推荐

