@@ -190,7 +190,7 @@ def loss_fn(model, config, data, dropout_rng, params, is_train=True):
190190 return loss , aux
191191
192192
193- def train_step (model , config , state_mesh_shardings , state , data , dropout_rng ):
193+ def train_step (model , config , state_mesh_shardings , params_shardings , state , data , dropout_rng ):
194194 """
195195
196196 Args:
@@ -212,12 +212,27 @@ def train_step(model, config, state_mesh_shardings, state, data, dropout_rng):
212212 extra_dpo_args = [reference_params ]
213213 _loss_fn = dpo_loss_fn
214214
215+ params = state .params
216+
215217 if config .gradient_accumulation_steps > 1 :
218+ # When using Zero-1 optimizer sharding, cast params to lower precision and apply sharding constraints
219+ # so that all-gather is done once in the lower precision before the gradient accumulation loop
220+ if config .shard_optimizer_over_data :
221+ def convert_to_bf16 (param ):
222+ if param .dtype == jnp .float32 :
223+ return param .astype (jnp .bfloat16 )
224+ else :
225+ return param
226+ ga_params = jax .tree_util .tree_map (convert_to_bf16 , params )
227+ ga_params = jax .tree .map (jax .lax .with_sharding_constraint , ga_params , params_shardings )
228+ else :
229+ ga_params = params
216230
217231 def accumulate_gradient (acc_grad_and_loss , data ):
232+ ga_params = acc_grad_and_loss ["ga_params" ]
218233 grad_func = jax .value_and_grad (_loss_fn , argnums = 4 , has_aux = True )
219234 (_ , aux ), cur_batch_gradient = grad_func (
220- model , config , data , dropout_rng , state . params , * extra_dpo_args , is_train = True
235+ model , config , data , dropout_rng , ga_params , * extra_dpo_args , is_train = True
221236 )
222237 acc_grad_and_loss ["loss" ] += aux ["total_loss" ]
223238 acc_grad_and_loss ["moe_lb_loss" ] += aux ["moe_lb_loss" ]
@@ -235,8 +250,16 @@ def reshape_to_microbatch_accumulations(batch_arr):
235250 return jnp .reshape (batch_arr , microbatch_shape )
236251
237252 data = jax .tree_util .tree_map (reshape_to_microbatch_accumulations , data )
238- init_grad = jax .tree_util .tree_map (jnp .zeros_like , state .params )
239- init_grad_and_loss = {"loss" : 0.0 , "grad" : init_grad , "total_weights" : 0 , "moe_lb_loss" : 0.0 , "mtp_loss" : 0.0 }
253+ init_grad = jax .tree_util .tree_map (jnp .zeros_like , ga_params )
254+ init_grad = jax .tree .map (jax .lax .with_sharding_constraint , init_grad , params_shardings )
255+ init_grad_and_loss = {
256+ "loss" : 0.0 ,
257+ "grad" : init_grad ,
258+ "total_weights" : 0 ,
259+ "moe_lb_loss" : 0.0 ,
260+ "mtp_loss" : 0.0 ,
261+ "ga_params" : ga_params ,
262+ }
240263
241264 grad_and_loss , aux = jax .lax .scan (
242265 accumulate_gradient , init_grad_and_loss , data , length = config .gradient_accumulation_steps
@@ -246,7 +269,10 @@ def reshape_to_microbatch_accumulations(batch_arr):
246269 + grad_and_loss ["moe_lb_loss" ] / config .gradient_accumulation_steps
247270 + grad_and_loss ["mtp_loss" ] / config .gradient_accumulation_steps
248271 )
249- raw_grads = jax .tree_util .tree_map (lambda arr : arr / grad_and_loss ["total_weights" ], grad_and_loss ["grad" ])
272+ raw_grads = grad_and_loss ["grad" ]
273+ if config .shard_optimizer_over_data :
274+ raw_grads = jax .tree .map (jax .lax .with_sharding_constraint , raw_grads , params_shardings )
275+ raw_grads = jax .tree_util .tree_map (lambda arr : arr / grad_and_loss ["total_weights" ], raw_grads )
250276 aux = jax .tree .map (lambda x : jnp .sum (x , axis = 0 ), aux ) # pytype: disable=module-attr
251277 else :
252278 if config .optimizer_memory_host_offload :
@@ -255,8 +281,10 @@ def reshape_to_microbatch_accumulations(batch_arr):
255281 reference_params , max_utils .with_memory_kind (reference_params_sharding , "device" )
256282 )
257283 extra_dpo_args = [reference_params ]
284+ if config .shard_optimizer_over_data :
285+ params = jax .tree .map (jax .lax .with_sharding_constraint , params , params_shardings )
258286 grad_func = jax .value_and_grad (_loss_fn , argnums = 4 , has_aux = True )
259- (loss , aux ), raw_grads = grad_func (model , config , data , dropout_rng , state . params , * extra_dpo_args , is_train = True )
287+ (loss , aux ), raw_grads = grad_func (model , config , data , dropout_rng , params , * extra_dpo_args , is_train = True )
260288
261289 raw_grads = jax .tree_util .tree_map (lambda x : x .astype (config .grad_dtype ) if x .dtype == jnp .float32 else x , raw_grads )
262290 intermediate_outputs = aux ["intermediate_outputs" ]
@@ -373,12 +401,15 @@ def train_loop(config, recorder, state=None):
373401 state = _merge_dpo_state (state , reference_params )
374402 state_mesh_shardings = _merge_dpo_state (state_mesh_shardings , state_mesh_shardings .params ["params" ])
375403
404+ params_shardings , state_mesh_shardings = maxtext_utils .maybe_update_params_sharding_with_opt (config , state_mesh_shardings )
405+
376406 p_train_step , p_eval_step = train_utils .jit_train_and_eval_step (
377- config , model , mesh , state , state_mesh_shardings , train_step , eval_step , eval_data_iterator
407+ config , model , mesh , state , state_mesh_shardings , train_step , eval_step , eval_data_iterator , params_shardings
378408 )
379409
380410 with mesh , nn_partitioning .axis_rules (config .logical_axis_rules ):
381411 shaped_batch = maxtext_utils .get_shaped_batch (config )
412+ state = jax .lax .with_sharding_constraint (state , state_mesh_shardings )
382413 compiled = p_train_step .lower (state , shaped_batch , init_rng ).compile ()
383414 compiled_stats = compiled .memory_analysis ()
384415 max_utils .print_compiled_memory_stats (compiled_stats )
@@ -402,6 +433,7 @@ def train_loop(config, recorder, state=None):
402433 nextrng = jax .jit (jax .random .fold_in )(init_rng , step )
403434 with maybe_record_goodput (recorder , GoodputEvent .STEP , step ):
404435 with mesh , nn_partitioning .axis_rules (config .logical_axis_rules ):
436+ state = jax .lax .with_sharding_constraint (state , state_mesh_shardings )
405437 state , metrics = p_train_step (state , example_batch , nextrng )
406438
407439 step_time_delta = datetime .datetime .now () - last_step_completion
@@ -474,9 +506,9 @@ def initialize(argv: Sequence[str]) -> tuple[pyconfig.HyperParameters, Any, Any]
474506 # TODO: mazumdera@ : ensure missing mandatory fields in base.yml are filled in in argv,
475507 # or fill in here
476508 config = pyconfig .initialize (argv )
477- jax .config .update ("jax_use_shardy_partitioner" , config .shardy )
478509 max_utils .print_system_information ()
479510 validate_train_config (config )
511+ jax .config .update ("jax_use_shardy_partitioner" , config .shardy )
480512 os .environ ["TFDS_DATA_DIR" ] = config .dataset_path or ""
481513 vertex_tensorboard_manager = VertexTensorboardManager ()
482514 if config .use_vertex_tensorboard or os .environ .get ("UPLOAD_DATA_TO_TENSORBOARD" ):
0 commit comments