运行Deep-Virtual-Try-on代码遇Adam.get_updates参数错误求助
解决TensorFlow 2.11中Adam.get_updates()参数不匹配的问题
错误原因
原代码是针对TensorFlow 1.x/旧版Keras编写的,而TensorFlow 2.x中Keras优化器的get_updates()方法签名发生了变更:
- 旧版签名:
get_updates(params, grads, loss) - TF2.x签名:
get_updates(loss, params)
你传入的weightsD,[],loss_D对应旧版的三个参数,但TF2.x的方法只接受loss和params两个参数,因此触发参数数量不匹配的错误。
修复方案
方案1:适配TF2.x的get_updates()参数顺序
直接调整参数顺序,并移除多余的空梯度列表,同时注意TF2.x中lr参数已更名为learning_rate:
# 替换原错误代码行 training_updates = Adam(learning_rate=lrD, beta_1=0.5).get_updates(loss_D, weightsD)
方案2:改用TF2.x推荐的梯度带(GradientTape)训练模式
这是TF2.x的标准训练方式,更易维护且兼容性更好:
# 初始化优化器(注意使用learning_rate替代lr) optimizer_D = Adam(learning_rate=lrD, beta_1=0.5) # 使用GradientTape计算梯度并生成更新操作 with tf.GradientTape() as tape: # 保留原代码中loss_D的计算逻辑 loss_D = ... # 你的损失计算代码 grads = tape.gradient(loss_D, weightsD) training_updates = optimizer_D.apply_gradients(zip(grads, weightsD))
额外注意事项
- 若代码中还有其他TF1.x风格的API(如
tf.Session()、tf.placeholder()等),可能需要进一步适配TF2.x的执行模式; - 建议优先使用方案2,因为它更符合TF2.x的设计理念,也能避免更多旧API的兼容性问题。
内容的提问来源于stack exchange,提问作者Amanchi Kishorebabu
相关产品推荐
相关产品推荐

