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

Function train_step

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

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected