Split DPO state to separate reference parameters.
(state)
| 23 | |
| 24 | |
| 25 | def _split_dpo_state(state): |
| 26 | """Split DPO state to separate reference parameters.""" |
| 27 | reference_params = state.params["reference_params"] |
| 28 | new_state = state.replace(params={k: v for k, v in state.params.items() if k != "reference_params"}) |
| 29 | return new_state, reference_params |
| 30 | |
| 31 | |
| 32 | def dpo_loss_fn(model, config, data, dropout_rng, params, reference_params, is_train=True): |
no outgoing calls
no test coverage detected