1.6 变量:图上唯一可写的状态


文档摘要

1.6 变量:图上唯一可写的状态 本节摘要:tf.Variable 是计算图上唯一可写的状态容器,训练的一切落脚点——参数被优化器一轮轮写入、指标被累加、计数器被推进,都发生在变量里。本节讲清 Variable 与 Tensor 的本质差异(可写性、生命周期、跨调用存续)、assign 与值的更新时序、可训练属性 trainable 的意义,以及变量复用与初始化时机的三个坑。没有变量就没有学习,这一节是训练机制的地基。 本节能力清单 阅读完本节,你应当能够: 说清 tf.Variable 与 tf.Tensor 在可写性、生命周期、身份三方面的差异; 用 assign、assignadd 正确更新变量,并解释为什么直接下标赋值会失败;

1.6 变量:图上唯一可写的状态

本节摘要:tf.Variable 是计算图上唯一可写的状态容器,训练的一切落脚点——参数被优化器一轮轮写入、指标被累加、计数器被推进,都发生在变量里。本节讲清 Variable 与 Tensor 的本质差异(可写性、生命周期、跨调用存续)、assign 与值的更新时序、可训练属性 trainable 的意义,以及变量复用与初始化时机的三个坑。没有变量就没有学习,这一节是训练机制的地基。

本节能力清单

阅读完本节,你应当能够:

  1. 说清 tf.Variable 与 tf.Tensor 在可写性、生命周期、身份三方面的差异;
  2. 用 assign、assign_add 正确更新变量,并解释为什么直接下标赋值会失败;
  3. 解释 trainable 属性如何决定一个变量是否进入求导名单;
  4. 识别"变量在循环外误建""重复创建未复用"两类典型错误。

为什么必须是变量

张量在排程室的规矩里是不可变的:一个 tf.Tensor 被创建后,值就定格了。这对数据流很友好——边上的值不会中途变卦,图的可读性、可优化性都建立在这条规矩上。但训练要求矛盾的事:参数必须在每一步被改写。框架的解法是给图开一个特批的"可写窗口"——变量。变量的值存在进程的设备内存里,图上的算子可以读它,优化器算完梯度后可以写它。一句话:Tensor 是边上的货物,Variable 是货架上那格可反复补货的位置。

import tensorflow as tf t = tf.constant([1.0, 2.0]) # t[0].assign(9.0) # 取消注释即报错:Tensor 不可写 # AttributeError: 'EagerTensor' object has no attribute 'assign' v = tf.Variable([1.0, 2.0]) # 创建变量:登记 + 分配可写内存 v[0].assign(9.0) # 合法:原地写入 print(v) # 输出:<tf.Variable 'Variable:0' shape=(2,) dtype=float32, # numpy=array([9., 2.], dtype=float32)> v.assign_add([1.0, 1.0]) # 原地加 print(v.numpy()) # 输出:[10. 3.]

可写性的边界:哪些写法有效

变量的写操作有明确清单:assign(整体覆盖)、assign_addassign_sub、以及切片 assign。所有写法都是原子的原地操作,返回的引用指向同一块内存:

w = tf.Variable([[1.0, 2.0], [3.0, 4.0]]) w.assign([[0.0, 0.0], [0.0, 0.0]]) # 整体覆盖 print(w.numpy()) # 输出:[[0. 0.] # [0. 0.]] w[0].assign([5.0, 6.0]) # 按行切片写 print(w.numpy()) # 输出:[[5. 6.] # [0. 0.]] # 以下写法都不会改 w 的值,初学者高发: # w = w + 1 —— 这是新建了一个 Tensor 并让名字 w 指向它,原变量没动 # w.numpy()[0][0] = 1 —— numpy() 返回的是副本,改副本无关原变量

w = w + 1 这个坑值得单独念一遍:它不报错、看起来"更新了",实际上 w 名字此后指向一个普通张量,原变量成了孤儿。Keras 训练能正常更新参数,正是因为优化器内部用的是 assign 类操作而非重新赋值。

trainable:进不进求导名单

变量创建时默认 trainable=True,会被自动收集进全局可训练变量集合,GradientTape 默认只对这些变量求导。像批次计数这类"要写但不要求导"的状态,应显式关掉:

weights = tf.Variable(tf.random.normal([3, 2]), name="weights") # 默认 trainable=True steps = tf.Variable(0, trainable=False, dtype=tf.int64) # 计数器 print([v.name for v in tf.trainable_variables()]) # 输出(示例):['weights:0'] —— steps 不在名单里 steps.assign_add(1) print(steps.numpy()) # 输出:1

tf.trainable_variables() 这份名单,正是 Keras 里 model.trainable_weights 的底层来源,也是优化器决定"我要更新谁"的依据。第 5 章自定义训练循环时,你会亲手把它交给优化器。

生命周期:变量死在谁手里

变量由 Python 对象持有,对象被回收、变量内存才释放。两个高频陷阱都在生命周期上。陷阱一,在循环或函数里反复创建变量——每次都新分配内存并新起名字:

# 错误示范:循环里建变量 def bad_stack(): parts = [] for i in range(3): v = tf.Variable(float(i), name="acc") # 三次创建:acc、acc_1、acc_2 parts.append(v) return parts print([p.name for p in bad_stack()]) # 输出:['acc:0', 'acc_1:0', 'acc_2:0'] —— 三个独立变量 # 正确做法:建一次,循环里 assign def good_stack(): acc = tf.Variable(0.0, name="acc") for i in range(3): acc.assign(acc.read_value() + float(i)) return acc.numpy() print(good_stack()) # 输出:3.0

陷阱二,在 @tf.function 里创建变量。跟踪只发生一次,变量也随之只创建一次,看似正常;但若两次跟踪(签名变了),同名变量重建会直接报错。规范是:变量在图函数外、模型层内创建,函数体只读只写不建。Keras 层类帮你把这个规矩封装好了,每个 Layer 的 build 方法只会为参数建一次变量。

class Counter(tf.Module): def __init__(self): # 在构造器里建变量:整个对象生命周期只建一次 self.count = tf.Variable(0.0) @tf.function def bump(self): self.count.assign_add(1.0) return self.count c = Counter() print(c.bump().numpy()) # 1.0 print(c.bump().numpy()) # 2.0 —— 同一个变量被持续累加

⚠️ 常见坑:在训练循环里用 w = w - lr * grad 形式"更新"参数。这不改写变量,只是让 Python 名字改指新张量,下一轮梯度对着新张量算,而优化器与 checkpoint 还认旧变量——整条训练链路静默断开。更新参数只用 assign 或优化器。

💡 关键直觉:看到"状态跨步保存"的需求就想到 Variable——计数器、累计指标、滑动平均、动量,全是变量。看到"每步都要的新值"就想到普通张量。分清这两类,图的内存模型就清晰了。

本节要点回顾

  • Variable 是唯一可写状态:Tensor 定格在边上,Variable 是货架上可补货的位置。
  • 写入只用 assign 族w = w + 1 是换名字不是改值,是训练静默失效的头号原因。
  • trainable 决定求导名单:计数器、计数类状态显式 trainable=False。
  • 变量建在循环外:循环内建变量造成内存泄漏与命名漂移。
  • 变量建在图函数外:模型层构造器或 build 方法里创建,跟踪期间不重复建。

下节把全册用到的模块收进一张地图。


作者与出处
原作者: 灏天文库
来源:灏天文库
整理: 灏天文库整理
由灏天文库平台收录,内容或由平台用户上传,仅供学习交流
发布者: 作者: 灏天文库 转发
评论区 (0)
U