(
model: _model.BaseModel, rng: at.KeyArrayLike, observation: _model.Observation, actions: _model.Actions
)
| 145 | |
| 146 | @at.typecheck |
| 147 | def loss_fn( |
| 148 | model: _model.BaseModel, rng: at.KeyArrayLike, observation: _model.Observation, actions: _model.Actions |
| 149 | ): |
| 150 | chunked_loss = model.compute_loss(rng, observation, actions, train=True) |
| 151 | return jnp.mean(chunked_loss) |
| 152 | |
| 153 | train_rng = jax.random.fold_in(rng, state.step) |
| 154 | observation, actions = batch |
nothing calls this directly
no test coverage detected