aopstudio 的个人博客

记录精彩的程序人生

AOP=art of programming=编程的艺术=程艺
  menu
75 文章
0 浏览
4 当前访客
ღゝ◡╹)ノ❤️

Python 中容易被忽略的原地操作:从 PyTorch Tensor 的 `-=` 说起

前言

最近重新开始学习《动手学深度学习》(D2L)。在从零实现线性回归时,我看到了这样一段小批量随机梯度下降代码:

def sgd(params, lr, batch_size):
    with torch.no_grad():
        for param in params:
            param -= lr * param.grad / batch_size
            param.grad.zero_()

一开始让我感到奇怪的是这一行:

param -= lr * param.grad / batch_size

param 明明只是函数里的一个局部变量,为什么修改它之后,函数外面的 wb 也跟着发生了变化?这个函数甚至不需要 return w, b,模型参数就已经被更新了。

继续研究以后,我才注意到这里涉及 Python 中一个很基础、却很容易被忽略的知识点:原地操作(in-place operation)。所谓原地操作,就是直接修改已有对象的内部状态,而不是创建一个新对象。

不过,只知道这个定义还不足以解释开头的问题。我们还需要弄清楚:函数内外的变量是否指向同一个对象,以及函数中的代码究竟修改了对象本身,还是只改变了局部变量的指向。

这个概念其实属于 Python 的基础知识。但我初学 Python 时并没有真正留意,后来写代码时也很少停下来想清楚。直到这次在 PyTorch 中看到一个看似普通的 -=,才发现自己需要回来补上这一课。

要把这件事说清楚,不妨先暂时离开 Tensor,从一个最普通的 Python 函数开始。

先看一个没有影响外部的例子

下面这个函数试图给传入的数字加一:

def add_one(x):
    x += 1

a = 1
add_one(a)

print(a)
# 1

函数中的 x 明明执行了 x += 1,为什么外面的 a 仍然是 1

要回答这个问题,首先要把变量对象分开。执行 a = 1 时,可以简单理解为变量 a 指向整数对象 1

a ───> 1

调用 add_one(a) 时,函数内部的形参 x 也会绑定到这个对象:

外部 a ──┐
         ├──> 1
函数 x ──┘

但是,整数是不可变对象(immutable object)。对象 1 创建以后,不能被直接修改成 2。因此,x += 1 会得到一个新的整数对象 2,再让函数内的局部变量 x 指向它:

外部 a ───> 1

函数 x ───> 2

改变的只是局部变量 x 的指向,外部的 a 从头到尾都没有换过对象,所以它仍然是 1。这种让变量改为指向另一个对象的过程,通常称为重新绑定(rebinding)

Python 中常见的不可变对象还包括 floatboolstrtuple。例如:

s = "hello"
s += " world"

这不是把原字符串直接延长,而是创建一个新字符串,再让 s 指向它。tuple 本身同样不可变,不过它可以包含列表之类的可变对象;修改元组中的列表,并不等于改变了元组自身保存的元素位置。

看到这里,可能很容易得出一个过于简单的结论:函数里的变量只是局部变量,所以修改它不会影响外部。可是把整数换成列表,结果就不一样了。

换成列表,为什么外部也变了?

来看下面这个函数:

def append_one(x):
    x.append(3)

a = [1, 2]
append_one(a)

print(a)
# [1, 2, 3]

这一次,函数没有返回任何内容,外部的列表却发生了变化。

调用 append_one(a) 时,外部的 a 和函数内部的 x 指向同一个列表对象:

外部 a ──┐
         ├──> [1, 2]
函数 x ──┘

列表与整数不同,它是可变对象(mutable object)。对象创建以后,其中的内容仍然可以改变。x.append(3) 没有让 x 指向另一个列表,而是直接修改了 ax 共同指向的对象:

外部 a ──┐
         ├──> [1, 2, 3]
函数 x ──┘

所以,虽然代码中写的是 x.append(3),外部通过 a 也能看到变化。这里不是函数把结果“传回”了外部,而是函数内外一直在观察同一个对象。

Python 中常见的可变对象还包括 dictset。例如:

d = {"a": 1}
d["b"] = 2

s = {1, 2}
s.add(3)

这些操作改变的也都是已有对象的内容。

总结一下:调用函数时,实参所指向的对象会绑定给形参;如果函数修改了这个对象,外部就能看到变化;如果只是让形参指向另一个对象,外部变量的绑定就不会跟着改变。

不过,这里还有一个问题:只要传入的是可变对象,函数里的操作就一定会影响外部吗?

可变对象就一定会影响外部吗?

仍然使用列表,但把函数改成普通加法:

def append_one(x):
    x = x + [3]

a = [1, 2]
append_one(a)

print(a)
# [1, 2]

这一次,传入的仍然是可变对象,外部的 a 却没有变化。

原因在于,x + [3] 会创建一个新列表 [1, 2, 3],随后 x = ... 只是让函数中的局部变量指向这个新列表。原来的列表没有被修改:

外部 a ───> [1, 2]

函数 x ───> [1, 2, 3]

这说明“对象是否可变”只回答了它能不能被直接修改,却不能告诉我们某一行代码有没有这样做。判断外部是否会受到影响,还要区分修改对象和重新绑定变量。

如果把函数中的普通加法换成增强赋值,结果又会不同:

def append_one(x):
    x += [3]

a = [1, 2]
append_one(a)

print(a)
# [1, 2, 3]

对于 list 来说,+= 会直接扩展原列表,也就是原地操作。函数中的 x 没有换对象,改变的是 ax 共同指向的列表:

外部 a ──┐
         ├──> [1, 2, 3]
函数 x ──┘

因此,下面两种写法虽然都能让函数内的 x 得到 [1, 2, 3],背后发生的事情并不相同:

x = x + [3]  # 创建新列表,再重新绑定 x
x += [3]     # 对 list 来说,会修改原列表

这里还要注意一个边界:+=-= 这些增强赋值运算符并不天然等于原地操作,具体行为仍然取决于对象的类型。对于 list+= 可以修改原列表;对于 intstr 等不可变对象,它只能产生新对象,再重新绑定变量。

到这里,判断这类问题所需要的两层关系已经比较清楚了:

  1. 函数内外的变量是否指向同一个对象?
  2. 函数中的代码修改了这个对象,还是让局部变量重新绑定到了另一个对象?

带着这两个问题,现在可以回到最开始的 PyTorch 代码了。

回到 PyTorch 中的 Tensor

再看 D2L 中的随机梯度下降函数:

def sgd(params, lr, batch_size):
    with torch.no_grad():
        for param in params:
            param -= lr * param.grad / batch_size
            param.grad.zero_()

调用函数时传入:

sgd([w, b], lr, batch_size)

列表 params 中保存的不是 wb 的独立副本,而是它们所指向的 Tensor 对象。执行 for param in params 时,局部变量 param 会先绑定到 w 对应的 Tensor,再绑定到 b 对应的 Tensor。

因此:

param -= lr * param.grad / batch_size

会原地修改当前 Tensor 中的数据。函数外面的 wb 仍然指向这些 Tensor,所以也能看到更新后的参数。这就是 sgd 不需要返回 wb 的原因。

如果改成:

param = param - lr * param.grad / batch_size

右侧会产生一个新的 Tensor,随后只有函数中的局部变量 param 改为指向它。外部的 w 不会因此重新绑定到这个新 Tensor:

外部 w ───> 原 Tensor

param ────> 新 Tensor

这与前面列表中的 x = x + [3] 本质上是同一件事。PyTorch 并没有改变 Python 中变量与对象的关系,只是 Tensor 的 -= 在这里执行了原地操作。

扩展:PyTorch 中的 _ 方法与自动求导

除了增强赋值,PyTorch 中还经常看到以下划线结尾的 Tensor 方法:

zero_()
add_()
sub_()
mul_()
copy_()

在常见的 Tensor API 中,方法名末尾的 _ 表示它会原地修改 Tensor。例如:

x.zero_()

这不是创建一个新的全零 Tensor,而是直接把 x 中的数据清零。D2L 代码中的:

param.grad.zero_()

也是在原地清空当前的梯度 Tensor。

不过,Tensor 与普通 Python 可变对象相比,还多了自动求导这一层问题。autograd 在反向传播时可能需要使用前向计算中保存的值;如果这些值在反向传播之前被原地覆盖,梯度计算就可能受到影响。因此,对参与计算图的 Tensor 进行原地操作需要格外小心,PyTorch 也会对可能破坏梯度计算的情况进行检查。

D2L 在更新参数时专门写了:

with torch.no_grad():
    param -= ...

参数更新本身不需要被 autograd 记录,所以要放在 torch.no_grad() 环境中执行。至于计算图如何记录操作、为什么有些原地修改会报错,属于自动求导机制的另一个话题,本文先不展开。

最后再看一个容易混淆的情况

开头的 param -= ... 已经可以解释清楚了。不过,还有一种写法很适合用来检查前面的判断方法是否真的掌握了:

def add_one(x):
    for i in range(len(x)):
        x[i] = x[i] + 1

x = [1, 2]
add_one(x)

print(x)
# [2, 3]

这里没有使用 +=,右侧的加法还产生了新的整数对象,为什么函数外面的列表仍然改变了?

关键在于赋值语句的左边:

x[i] = x[i] + 1

右侧的 x[i] + 1 确实会产生一个新整数,但左侧并不是让变量 x 重新绑定,而是把结果写回 x 所指向列表的第 i 个位置。

以第一次循环为例,原来的列表是 [1, 2]。右侧先得到整数 2,随后左侧把列表中的第一个元素替换成它,于是列表变成 [2, 2]。第二次循环后,列表再变成 [2, 3]

所以,“这算不算原地操作”需要先说清楚讨论的是哪个对象:从整数对象的角度看,原来的整数没有被修改,而是被一个新整数替换;从列表对象的角度看,x 始终指向原来的列表,变化发生在列表内部,因此这是对列表的原地修改。

更准确地说,x[i] = ... 是一次下标赋值,它的效果是在原列表上替换一个元素,而不是修改那个不可变的整数对象。判断一行代码会不会影响外部,不能只看有没有 +=,也不能只看右侧是否创建了新对象,还要看左侧的赋值目标:

x = ...       # 让变量 x 重新绑定
x += ...      # 可能原地修改对象,取决于具体类型
x[i] = ...    # 对列表来说,会修改原列表中的一个位置
x.attr = ...  # 设置 x 所指向对象的属性

可变与不可变决定的是对象是否允许自身状态发生变化;具体的一行代码是否进行了原地修改,还要结合操作本身、赋值目标以及正在讨论的对象层级来判断。

总结

回头看,D2L 中的这一行代码:

param -= lr * param.grad / batch_size

之所以能够在没有返回值的情况下更新函数外面的 wb,需要同时满足两个条件:局部变量 param 与外部变量指向同一个 Tensor,而 -= 修改的是这个 Tensor 本身,不只是让 param 指向一个新对象。

以后再遇到类似问题,可以依次问自己:

  1. 这些变量是否指向同一个对象?
  2. 这个对象是可变对象还是不可变对象?
  3. 当前操作修改了对象,还是重新绑定了变量?
  4. 如果存在赋值,左侧的目标究竟是变量、对象中的位置,还是对象的属性?

这几个问题比单纯记住“可变对象会变”或者“+= 是原地操作”更可靠。


标题:Python 中容易被忽略的原地操作:从 PyTorch Tensor 的 `-=` 说起
作者:aopstudio
地址:https://neusoftware.top/articles/2026/09/09/1788949569832.html