(model_params: Nested[Tensor], *, inputs: Any)
| 669 | model_parameters_grad = jax.tree.map( |
| 670 | lambda compute_gradients, v: v if compute_gradients else dummy_value, |
| 671 | should_compute_gradients, |
| 672 | model_params, |
| 673 | ) |
| 674 | model_parameters_no_grad = jax.tree.map( |
| 675 | lambda compute_gradients, v: dummy_value if compute_gradients else v, |
| 676 | should_compute_gradients, |
| 677 | model_params, |
| 678 | ) |
| 679 | return model_parameters_grad, model_parameters_no_grad |
| 680 | |
| 681 | return filtered_forward, split_params_fn |
| 682 |