You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.26 07:16:29