MCPcopy Create free account
hub / github.com/LargeWorldModel/LWM / eval_step

Function eval_step

lwm/train.py:225–270  ·  view source on GitHub ↗
(train_state, rng, batch)

Source from the content-addressed store, hash-verified

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(

Callers

nothing calls this directly

Calls 1

rng_keysMethod · 0.80

Tested by

no test coverage detected