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

TensorFlow自定义优化器:如何将参数与梯度转为向量并更新

解决TensorFlow自定义优化器中参数/梯度的向量转换与更新问题

问题原因

TensorFlow的Optimizer默认按单个参数张量(比如卷积层的kernel、偏置bias)处理梯度更新,而非将所有参数拼接成一个全局向量。这就是你看到grad和var是多维张量而非一维向量的原因。要实现基于全局向量的优化算法,需要手动完成「全局向量拼接→算法更新→拆分回原张量」的流程。

解决方案步骤

  1. 收集并拼接全局参数与梯度:遍历所有可训练变量,将每个变量和对应的梯度flatten后拼接成一维全局向量x(参数)和g(梯度)。
  2. 执行自定义优化更新:用你的算法逻辑更新全局向量x得到x_new。
  3. 拆分并更新原变量:将x_new按原变量的形状拆分,逐个赋值回对应的参数张量。

修改后的完整代码

import tensorflow as tf
import numpy as np

class TestGD(tf.keras.optimizers.Optimizer):
    def __init__(self, rad=0.01, learning_rate=0.01,
                 use_locking=False, name="TestGD"):
        super(TestGD, self).__init__(use_locking, name)
        self._radius = rad
        self._learning_rate = learning_rate

    def _create_slots(self, var_list):
        # 可在此创建优化器需要维护的状态变量(如动量项)
        pass

    def _prepare(self):
        self._learning_rate_t = tf.convert_to_tensor(self._learning_rate, name="learning_rate")
        self._radius_t = tf.convert_to_tensor(self._radius, name="radius")

    def apply_gradients(self, grads_and_vars, name=None):
        # 过滤无梯度的变量
        grads_and_vars = [(g, v) for g, v in grads_and_vars if g is not None]
        if not grads_and_vars:
            return tf.no_op(name=name)

        # 1. 拼接全局参数向量x和梯度向量g
        vars_flat = []
        grads_flat = []
        shapes = []
        for grad, var in grads_and_vars:
            shapes.append(var.shape)
            vars_flat.append(tf.reshape(var, [-1]))
            grads_flat.append(tf.reshape(grad, [-1]))
        
        x = tf.concat(vars_flat, axis=0)
        g = tf.concat(grads_flat, axis=0)

        # 2. 替换为你的自定义优化算法逻辑
        # 示例:带范数约束的梯度下降
        x_new = x - self._learning_rate_t * g
        # 可选:添加参数范数约束
        norm = tf.norm(x_new)
        x_new = tf.cond(norm > self._radius_t, 
                        lambda: x_new * self._radius_t / norm,
                        lambda: x_new)

        # 3. 拆分更新后的向量,逐个赋值回原变量
        splits = tf.split(x_new, [tf.reduce_prod(shape) for shape in shapes])
        updates = []
        for split, (grad, var) in zip(splits, grads_and_vars):
            updated_var = tf.reshape(split, var.shape)
            updates.append(tf.assign(var, updated_var))

        return tf.group(*updates, name=name)

    def _apply_dense(self, grad, var):
        # 已重写apply_gradients,该方法无需实现
        raise NotImplementedError("Use apply_gradients instead.")

    def _apply_sparse(self, grad, var):
        raise NotImplementedError("Sparse gradient updates are not supported.")

# Build LeNet model
model = tf.keras.Sequential([
    tf.keras.layers.Conv2D(6, kernel_size=(5, 5), activation='relu', input_shape=(28, 28, 1)),
    tf.keras.layers.MaxPooling2D(pool_size=(2, 2)),
    tf.keras.layers.Conv2D(16, kernel_size=(5, 5), activation='relu'),
    tf.keras.layers.MaxPooling2D(pool_size=(2, 2)),
    tf.keras.layers.Flatten(),
    tf.keras.layers.Dense(120, activation='relu'),
    tf.keras.layers.Dense(84, activation='relu'),
    tf.keras.layers.Dense(10, activation='softmax')
])

# 使用自定义优化器
custom_optimizer = TestGD(rad=1.0, learning_rate=0.01)

# 编译模型
model.compile(optimizer=custom_optimizer,
              loss='sparse_categorical_crossentropy',
              metrics=['accuracy'])

# 加载并处理数据集
(x_train, y_train), (x_test, y_test) = tf.keras.datasets.mnist.load_data()
x_train, x_test = x_train / 255.0, x_test / 255.0
x_train = x_train[..., tf.newaxis].astype("float32")
x_test = x_test[..., tf.newaxis].astype("float32")

train_dataset = tf.data.Dataset.from_tensor_slices((x_train, y_train)).shuffle(60000).batch(64)
test_dataset = tf.data.Dataset.from_tensor_slices((x_test, y_test)).batch(64)

# 训练与评估
model.fit(train_dataset, epochs=5)
test_loss, test_acc = model.evaluate(test_dataset)
print(f"Test accuracy: {test_acc}")

关键细节说明

  • 重写apply_gradients:这是实现全局向量操作的核心,该方法能获取所有(梯度, 变量)对,而_apply_dense仅处理单个变量。
  • 过滤无梯度变量:部分变量可能无需梯度(如冻结层),提前过滤可避免报错。
  • TensorFlow原生操作:全程使用tf.reshape、tf.concat等操作,确保兼容图模式和Eager Execution,避免用numpy操作破坏计算图。
  • 状态变量处理:若优化器需维护状态(如动量项),可在_create_slots中创建对应变量,同样通过拼接/拆分同步全局状态。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 04:37:03