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

如何在Keras中使用tensorflow_graphics的levenberg_marquardt优化器?

解决Keras中使用tensorflow_graphics的Levenberg-Marquardt优化器的问题

问题原因

你遇到的错误核心是:tensorflow_graphics中的levenberg_marquardt不是Keras兼容的优化器类,它本质是一个用于梯度优化的函数,而非tf.keras.optimizers.Optimizer的子类实例,所以无法直接传给model.compile()的optimizer参数。直接引用模块(不加括号)或尝试实例化(加括号)都不符合Keras对优化器的要求。

解决方法

需要通过自定义训练循环来调用TFG的Levenberg-Marquardt优化函数,手动完成模型参数的更新,而非依赖Keras的compile()和fit()流程。

修改后的完整代码示例

from keras.models import Sequential
import keras
import tensorflow as tf
import tensorflow_graphics as tfg

# 构建模型
model = keras.Sequential([keras.layers.Dense(3, activation=tf.nn.relu, input_shape=[4])])

# 定义损失函数
def loss_fn(model, x, y):
    y_pred = model(x, training=True)
    return tf.reduce_mean(tf.keras.losses.mean_squared_error(y, y_pred))

# 准备示例训练数据
x_train = tf.random.normal((100, 4))
y_train = tf.random.normal((100, 3))

# 初始化TFG的Levenberg-Marquardt优化状态
optimizer_state = tfg.math.optimizer.levenberg_marquardt.initialize(model.trainable_variables)

# 自定义训练循环
epochs = 10
for epoch in range(epochs):
    print(f"Epoch {epoch+1}/{epochs}")
    # 调用TFG优化器计算损失并更新参数
    loss, optimizer_state, updated_vars = tfg.math.optimizer.levenberg_marquardt.minimize(
        loss_fn,
        variables=model.trainable_variables,
        var_list=model.trainable_variables,
        initial_optimizer_state=optimizer_state,
        x=x_train,
        y=y_train
    )
    # 将更新后的参数赋值回模型
    for var, updated_var in zip(model.trainable_variables, updated_vars):
        var.assign(updated_var)
    # 打印当前训练损失
    print(f"Loss: {loss.numpy():.4f}")

关键说明

  • TFG的levenberg_marquardt.minimize()函数需要接收损失函数、待优化的变量列表、初始优化状态,以及损失函数所需的输入参数(如示例中的x和y)。
  • 优化完成后必须手动将更新后的参数赋值回模型的可训练变量,才能让模型应用新参数。
  • 这种方式绕开了Keras标准优化器的适配要求,直接利用TFG的优化逻辑完成参数更新。

内容的提问来源于stack exchange,提问作者Marc Fischer

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.04 07:15:29