| 223 | return train_state, rng_generator(), metrics |
| 224 | |
| 225 | def eval_step(train_state, rng, batch): |
| 226 | rng_generator = JaxRNG(rng) |
| 227 | batch = with_sharding_constraint(batch, PS(('dp', 'fsdp'), 'sp')) |
| 228 | if FLAGS.modality == 'text': |
| 229 | logits = model.apply( |
| 230 | train_state.params, |
| 231 | batch['input_tokens'], |
| 232 | deterministic=True, |
| 233 | rngs=rng_generator(llama_config.rng_keys()), |
| 234 | ).logits |
| 235 | loss, acc = cross_entropy_loss_and_accuracy( |
| 236 | logits, |
| 237 | batch['target_tokens'], |
| 238 | batch['loss_masks'] |
| 239 | ) |
| 240 | metrics = dict( |
| 241 | eval_loss=loss, |
| 242 | eval_acc=acc, |
| 243 | ) |
| 244 | elif FLAGS.modality == 'vision,text': |
| 245 | vision_logits, text_logits = model.apply( |
| 246 | train_state.params, |
| 247 | batch['input_tokens'], |
| 248 | batch['input_vision_masks'], |
| 249 | deterministic=True, |
| 250 | rngs=rng_generator(llama_config.rng_keys()), |
| 251 | ).logits |
| 252 | vision_loss, vision_acc = cross_entropy_loss_and_accuracy( |
| 253 | vision_logits, |
| 254 | jnp.where(batch['target_vision_masks'], batch['target_tokens'], 0), |
| 255 | batch['loss_masks'] * batch['target_vision_masks'] |
| 256 | ) |
| 257 | text_loss, text_acc = cross_entropy_loss_and_accuracy( |
| 258 | text_logits, |
| 259 | jnp.where(batch['target_vision_masks'], 0, batch['target_tokens']), |
| 260 | batch['loss_masks'] * (1.0 - batch['target_vision_masks']) |
| 261 | ) |
| 262 | loss = 0.5 * (vision_loss + text_loss) |
| 263 | metrics = dict( |
| 264 | eval_loss=loss, |
| 265 | eval_vision_accuracy=vision_acc, |
| 266 | eval_vision_loss=vision_loss, |
| 267 | eval_text_accuracy=text_acc, |
| 268 | eval_text_loss=text_loss, |
| 269 | ) |
| 270 | return rng_generator(), metrics |
| 271 | |
| 272 | train_state_shapes = jax.eval_shape(init_fn, next_rng()) |
| 273 | train_state_partition = match_partition_rules( |