(g, p, s)
| 607 | param_scales = _weight_decay_scales(params, per_param_scale=per_param_scale) |
| 608 | |
| 609 | def f(g, p, s): |
| 610 | return g + weight_decay * lr_scale * p.value * s |
| 611 | |
| 612 | updates = jax.tree.map( |
| 613 | lambda x, y, z: None if x is None else f(x, y, z), |
no outgoing calls