(train_state, rng, batch)
| 164 | return TrainState.create(params=params, tx=optimizer, apply_fn=None) |
| 165 | |
| 166 | def train_step(train_state, rng, batch): |
| 167 | rng_generator = JaxRNG(rng) |
| 168 | batch = with_sharding_constraint(batch, PS(('dp', 'fsdp'), 'sp')) |
| 169 | def loss_and_accuracy(params): |
| 170 | if FLAGS.modality == 'text': |
| 171 | logits = model.apply( |
| 172 | params, |
| 173 | batch['input_tokens'], |
| 174 | deterministic=False, |
| 175 | rngs=rng_generator(llama_config.rng_keys()), |
| 176 | ).logits |
| 177 | loss, acc = cross_entropy_loss_and_accuracy( |
| 178 | logits, |
| 179 | batch['target_tokens'], |
| 180 | batch['loss_masks'] |
| 181 | ) |
| 182 | metrics = dict(acc=acc) |
| 183 | return loss, metrics |
| 184 | elif FLAGS.modality == 'vision,text': |
| 185 | vision_logits, text_logits = model.apply( |
| 186 | params, |
| 187 | batch['input_tokens'], |
| 188 | batch['input_vision_masks'], |
| 189 | deterministic=False, |
| 190 | rngs=rng_generator(llama_config.rng_keys()), |
| 191 | ).logits |
| 192 | vision_loss, vision_acc = cross_entropy_loss_and_accuracy( |
| 193 | vision_logits, |
| 194 | jnp.where(batch['target_vision_masks'], batch['target_tokens'], 0), |
| 195 | batch['loss_masks'] * batch['target_vision_masks'] |
| 196 | ) |
| 197 | text_loss, text_acc = cross_entropy_loss_and_accuracy( |
| 198 | text_logits, |
| 199 | jnp.where(batch['target_vision_masks'], 0, batch['target_tokens']), |
| 200 | batch['loss_masks'] * (1.0 - batch['target_vision_masks']) |
| 201 | ) |
| 202 | loss = 0.5 * (vision_loss + text_loss) |
| 203 | |
| 204 | metrics = dict( |
| 205 | vision_loss=vision_loss, |
| 206 | vision_acc=vision_acc, |
| 207 | text_loss=text_loss, |
| 208 | text_acc=text_acc, |
| 209 | ) |
| 210 | else: |
| 211 | raise ValueError(f"Unsupported modality: {FLAGS.modality}") |
| 212 | return loss, metrics |
| 213 | grad_fn = jax.value_and_grad(loss_and_accuracy, has_aux=True) |
| 214 | (loss, loss_metrics), grads = grad_fn(train_state.params) |
| 215 | train_state = train_state.apply_gradients(grads=grads) |
| 216 | metrics = dict( |
| 217 | loss=loss, |
| 218 | learning_rate=optimizer_info['learning_rate_schedule'](train_state.step), |
| 219 | param_norm=global_norm(train_state.params), |
| 220 | gradient_norm=global_norm(grads), |
| 221 | **loss_metrics |
| 222 | ) |
| 223 | return train_state, rng_generator(), metrics |
nothing calls this directly
no outgoing calls
no test coverage detected