TensorFlow中get_or_create_global_step属性错误及迁移求助
解决TensorFlow中
get_or_create_global_step的AttributeError问题 问题根源
tf.train.get_or_create_global_step()是TensorFlow 1.x的旧API,在TensorFlow 2.x中已被完全移除,调用时会触发AttributeError。你的代码还混用了其他TF1风格的API,需要同步迁移到TF2的写法。
分步修改方案
1. 手动创建全局步变量
替换原来的global_step = tf.train.get_or_create_global_step(),改为手动创建可追踪的TensorFlow变量:
# 初始化全局步为0,设置不可训练,命名为global_step global_step = tf.Variable(0, trainable=False, name='global_step')
如果需要在训练过程中更新步数,在每个训练步骤结束后执行global_step.assign_add(1)即可。
2. 替换优化器为TF2版本
把TF1的tf.train.AdamOptimizer换成TF2的tf.keras.optimizers.Adam:
generator_optimizer = tf.keras.optimizers.Adam(learning_rate) discriminator_optimizer = tf.keras.optimizers.Adam(learning_rate)
3. 替换tf.contrib.eager.defun为tf.function
tf.contrib.eager.defun是TF1 eager模式的旧装饰器,TF2中用tf.function替代:
# 替换原来的tf.contrib.eager.defun(train_step) train_step = tf.function(train_step)
更规范的写法是直接给train_step函数加装饰器:
@tf.function def train_step(...): # 你的训练逻辑代码 pass
4. 可选:将全局步加入检查点保存
如果需要把全局步和模型、优化器一起保存到检查点,修改tf.train.Checkpoint的初始化:
checkpoint = tf.train.Checkpoint( generator_optimizer=generator_optimizer, discriminator_optimizer=discriminator_optimizer, generator=generator, discriminator=discriminator, global_step=global_step # 添加全局步到检查点 )
修改后的关键代码片段
# Initialise logging log_path = os.path.join('logs', exp_name, time_string) summary_writer = tf.summary.create_file_writer(log_path, flush_millis=10000) summary_writer.set_as_default() # 替换全局步创建方式 global_step = tf.Variable(0, trainable=False, name='global_step') # ... 其余数据加载代码不变 ... # Set up the models for training generator = make_generator_model_small() discriminator = make_discriminator_model() # 替换优化器 generator_optimizer = tf.keras.optimizers.Adam(learning_rate) discriminator_optimizer = tf.keras.optimizers.Adam(learning_rate) checkpoint_prefix = os.path.join(model_path, "ckpt") # 可选:添加全局步到检查点 checkpoint = tf.train.Checkpoint(generator_optimizer=generator_optimizer, discriminator_optimizer=discriminator_optimizer, generator=generator, discriminator=discriminator, global_step=global_step) generate_and_save_images(None, 0, selected_inputs, selected_labels) # baseline print("\nTraining...\n") # 替换defun为tf.function train_step = tf.function(train_step) train(train_dataset, max_epoch) print("\nTraining done\n")
内容的提问来源于stack exchange,提问作者mchd
相关产品推荐
相关产品推荐

