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

在JAX中实现神经网络梯度自点积及Hessian-梯度计算

JAX/Flax 实现梯度与Hessian-梯度乘积的组合计算

问题描述

现有基于JAX的Flax模型训练代码:

(loss, (inner_state, logits)), grad = jax.value_and_grad(
    lambda m: forward_and_loss(m, true_gradient=True), has_aux=True)(model)

其中grad为flax.nn.Model类型。forward_and_loss函数定义如下:

def forward_and_loss(model: flax.nn.Model, true_gradient: bool = False):
    """Returns the model's loss, updated state and predictions.

    Args:
      model: The model that we are training.
      true_gradient: If true, the same mixing parameter will be used for the
        forward and backward pass for the Shake Shake and Shake Drop
        regularization (see papers for more details).
    """
    with flax.nn.stateful(state) as new_state:
        with flax.nn.stochastic(prng_key):
            try:
                logits = model(
                    batch['image'], train=True, true_gradient=true_gradient)
            except TypeError:
                logits = model(batch['image'], train=True)
    loss = cross_entropy_loss(logits, batch['label'])
    # We apply weight decay to all parameters, including bias and batch norm
    # parameters.
    weight_penalty_params = jax.tree_leaves(model.params)
    if FLAGS.no_weight_decay_on_bn:
        weight_l2 = sum(
            [jnp.sum(x**2) for x in weight_penalty_params if x.ndim > 1])
    else:
        weight_l2 = sum([jnp.sum(x ** 2) for x in weight_penalty_params])
    weight_penalty = l2_reg * 0.5 * weight_l2
    loss = loss + weight_penalty
    return loss, (new_state, logits)

需求:

  1. 计算grad与自身的点积(记为grad_grad)
  2. 求grad_grad关于模型参数的梯度,得到hessian_grad(即Hessian矩阵与grad的乘积)
  3. 最终计算grad + alpha * hessian_grad,且hessian_grad需保持flax.nn.Model类型。

实现代码

核心思路

  • 借助JAX的树结构工具处理Flax Model的嵌套参数,计算梯度点积
  • 通过二阶求导得到Hessian与梯度的乘积
  • 利用Flax Model的replace方法保持类型一致性

完整代码

import jax
import jax.numpy as jnp
import flax.nn

# 1. 计算grad与自身的点积
def grad_dot_product(grad):
    grad_leaves = jax.tree_util.tree_leaves(grad.params)
    return sum(jnp.sum(leaf ** 2) for leaf in grad_leaves)

grad_grad = grad_dot_product(grad)

# 2. 计算Hessian与grad的乘积(即grad_grad对模型参数的梯度)
def loss_grad_dot(model):
    (_, _), grad_model = jax.value_and_grad(
        lambda m: forward_and_loss(m, true_gradient=True), has_aux=True)(model)
    return grad_dot_product(grad_model)

# 得到参数结构的hessian梯度
hessian_grad_params = jax.grad(loss_grad_dot)(model)

# 3. 转换为flax.nn.Model类型
hessian_grad = grad.replace(params=hessian_grad_params.params)

# 4. 计算最终更新方向
alpha = 0.01  # 可根据需求调整系数
updated_grad_params = jax.tree_util.tree_map(
    lambda g, h: g + alpha * h, grad.params, hessian_grad.params
)
updated_grad_model = grad.replace(params=updated_grad_params)

关键说明

  • grad_dot_product通过jax.tree_util.tree_leaves提取所有参数叶子节点,计算平方和得到全局点积
  • loss_grad_dot封装梯度计算与点积流程,方便jax.grad直接求二阶导得到Hessian-梯度乘积
  • grad.replace(params=...)将参数结构映射回flax.nn.Model类型,保证与原grad类型一致
  • jax.tree_util.tree_map实现嵌套参数的逐元素运算,完成grad + alpha * hessian_grad的计算

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.17 21:37:30