如何更简洁地根据条件使用torch.no_grad()?
简洁化PyTorch中带条件的
torch.no_grad()写法 嘿,这个问题问到点子上了!重复写out=network(input)确实有点冗余,我给你几个简洁的解决方案,完美避开这段重复代码:
方法1:用contextlib.nullcontext()(最推荐)
这是最贴合你期望思路的写法,Python 3.7+自带的nullcontext()会在条件不满足时提供一个空上下文,完全不影响代码执行逻辑:
from contextlib import nullcontext # 根据条件选择对应的上下文管理器 ctx = torch.no_grad() if no_grad_condition else nullcontext() with ctx: out = network(input)
这样只需要写一次out=network(input),代码清晰又简洁,还完全符合PyTorch的使用习惯。
方法2:自定义带条件的上下文管理器
如果你想实现类似with torch.no_grad(no_grad_condition):的直观调用,可以自己封装一个小巧的上下文管理器:
import torch from contextlib import contextmanager @contextmanager def conditional_no_grad(enable): if enable: with torch.no_grad(): yield else: yield # 调用起来就和你想象的一样简洁 with conditional_no_grad(no_grad_condition): out = network(input)
这个自定义管理器可以在项目里重复使用,适合多处需要这个条件逻辑的场景。
方法3:一行表达式写法(不推荐)
如果非要追求极致行数,也可以用三元表达式结合上下文管理器的底层方法,但可读性会打折扣,不推荐在正式代码里用:
out = (torch.no_grad().__enter__(); res=network(input); torch.no_grad().__exit__(None,None,None); res) if no_grad_condition else network(input)
还是那句话,代码可读性永远是第一位的,前两种方法才是更优的选择。
内容的提问来源于stack exchange,提问作者Yuval Atzmon
相关产品推荐
相关产品推荐

