在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)
需求:
- 计算
grad与自身的点积(记为grad_grad) - 求
grad_grad关于模型参数的梯度,得到hessian_grad(即Hessian矩阵与grad的乘积) - 最终计算
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
相关产品推荐
相关产品推荐

