TensorFlow分布式策略下Keras自定义训练循环作用域异常咨询
问题解答
1. 优化器slot变量的创建时机与作用域保证
优化器的slot变量是用来存储动量、梯度累计值等优化过程状态的专属变量,创建时机为优化器第一次调用apply_gradients方法时,其所属的分布式策略作用域和调用apply_gradients时当前激活的策略上下文完全绑定,一旦创建完成就无法修改所属策略。
要保证slot变量作用域和模型变量一致,只需要遵循一个核心规则:模型权重、优化器实例的创建过程,以及首次apply_gradients调用,三者处于同一个strategy.scope()上下文下即可。
2. Sequential与函数式API的行为差异
二者的行为差异和模型结构无关,完全来自变量创建时机的底层实现不同:
- 函数式API在你执行
tf.keras.Model(inputs, outputs)实例化模型的瞬间,就会立刻创建所有层的权重变量,变量所属策略由实例化时的当前上下文决定。 - Sequential模型如果没有在初始化时传入
input_shape参数,会采用延迟初始化逻辑,直到第一次调用模型__call__方法喂入数据时才会真正创建权重变量,变量所属策略由第一次调用时的上下文决定。
所以当你不在strategy.scope()下创建模型时,函数式API的变量会直接绑定到默认单卡策略,而Sequential的变量还未创建,后续在分布式上下文下调用时才会生成对应策略下的变量,自然就会出现行为差异。
3. TensorFlow 2.4.1与2.5.0的行为差异原因
两个版本的结果相反来自TF团队对分布式策略集成逻辑的不兼容修改:
- 2.4版本存在隐式变量迁移逻辑:如果模型变量创建时不在分布式策略scope下,第一次在分布式上下文下调用模型时,Keras会自动将变量迁移到当前分布式策略下;同时该版本优化器slot创建逻辑存在bug,会自动继承模型变量的策略,忽略当前上下文,所以部分不符合规范的写法也能运行。
- 2.5版本修复了上述隐式迁移逻辑和slot创建bug,要求所有变量(模型权重、优化器slot)必须显式在对应分布式策略的scope下创建,不再支持自动迁移,不符合规范的写法会直接抛出作用域不一致的错误。
函数式API适配解决方案
推荐使用唯一兼容所有TF 2.x版本的标准写法,不需要修改模型结构:
将模型创建、优化器创建的全部逻辑都放在strategy.scope()上下文内,示例代码如下:
import tensorflow as tf # 初始化分布式策略 strategy = tf.distribute.MirroredStrategy() # 所有和模型、优化器相关的创建逻辑都放在scope下 with strategy.scope(): # 函数式API构建模型,和你原有业务逻辑完全一致 inputs = tf.keras.Input(shape=你的输入维度) x = 你的层逻辑(inputs) outputs = 输出层(x) model = tf.keras.Model(inputs, outputs) # 优化器也在scope下创建 optimizer = tf.keras.optimizers.你用的优化器() # 后续自定义训练循环逻辑和官方教程一致即可 @tf.function def train_step(dataset_inputs): def step_fn(inputs): x, y = inputs with tf.GradientTape() as tape: pred = model(x, training=True) loss = 你的损失函数(y, pred) grads = tape.gradient(loss, model.trainable_variables) optimizer.apply_gradients(zip(grads, model.trainable_variables)) return loss return strategy.run(step_fn, args=(dataset_inputs,))
如果因为业务逻辑限制无法将模型创建逻辑迁入scope,仅在TF 2.4下可以在第一次调用apply_gradients前手动触发模型变量迁移:先在strategy.scope()下调用一次模型的__call__方法喂入dummy数据,触发变量初始化后再执行训练逻辑,该方法不兼容TF 2.5及以上版本,不推荐长期使用。
内容的提问来源于stack exchange,提问作者Dan
相关产品推荐
相关产品推荐

