为何PyTorch张量在函数外可通过.data更新,却无法用+=操作?
PyTorch全局张量更新的两种方式差异解析
这本质是Python命名空间规则和PyTorch张量特性共同作用的结果,具体拆解如下:
1. x.data += ...能正常运行的原因
x.data是对全局变量x(张量对象)的属性访问,函数内部只是读取了全局作用域的x,然后对它的data属性执行原地修改——整个过程没有对x本身做赋值操作,Python会自动从全局命名空间查找x,因此不会触发局部变量的判定逻辑,自然能成功修改张量的底层数据。
代码示例:
x = torch.zeros(5) def my_function(): x.data += torch.ones(5) my_function() print(x) # tensor([1., 1., 1., 1., 1.])
2. x += ...报错的原因
Python里的+=属于增强赋值操作,对于PyTorch张量来说,它等价于x = x + torch.ones(5)——这本质是对变量x进行了赋值。当函数内部出现对变量的赋值行为时,Python会默认把该变量当作局部变量,但赋值语句右边的x在局部作用域里还未定义,因此会抛出UnboundLocalError。
错误代码示例:
x = torch.zeros(5) def my_function(): x += torch.ones(5) # UnboundLocalError: local variable 'x' referenced before assignment my_function()
解决x += ...报错的方法
如果想用+=更新全局张量,只需在函数内部用global声明x是全局变量,告诉Python不要将其视为局部变量:
x = torch.zeros(5) def my_function(): global x x += torch.ones(5) my_function() print(x) # tensor([1., 1., 1., 1., 1.])
内容的提问来源于stack exchange,提问作者jstm
相关产品推荐
相关产品推荐

