@@ -146,7 +146,7 @@ def restore_single_process(item, process_index, process_count):
146146
147147 if isinstance (item , list ):
148148 restored_items = []
149- for data_iter , process_idx in zip (item , process_index ):
149+ for data_iter , process_idx in zip (item , process_index ): # pyrefly: ignore[bad-argument-type]
150150 restored_items .append (restore_single_process (data_iter , process_idx , process_count ))
151151 return restored_items
152152 else :
@@ -423,7 +423,7 @@ def create_orbax_checkpoint_manager(
423423
424424 if dataset_type is not None and dataset_type == "grain" :
425425 item_names += ("iter" ,)
426- item_handlers ["iter" ] = GrainCheckpointHandler ()
426+ item_handlers ["iter" ] = GrainCheckpointHandler () # pyrefly: ignore[bad-assignment]
427427
428428 # local storage checkpoint needs parent directory created
429429 p = gcs_utils .mkdir_and_check_permissions (checkpoint_dir )
@@ -762,7 +762,7 @@ def load_state_if_possible(
762762 if checkpoint_manager is not None :
763763 max_logging .log ("checkpoint manager exists so trying to load this run's existing checkpoint" )
764764
765- step = checkpoint_manager .latest_step () if step < 0 else step
765+ step = checkpoint_manager .latest_step () if step < 0 else step # pyrefly: ignore[bad-assignment]
766766 if step is not None :
767767 max_logging .log (f"restoring from this run's directory step { step } " )
768768
@@ -1150,10 +1150,10 @@ def save_checkpoint(checkpoint_manager, step, state, config=None, data_iterator=
11501150 if config and config .dataset_type == "grain" and not isinstance (data_iterator , PlaceHolderDataIterator ):
11511151 if isinstance (data_iterator , RemoteIteratorWrapper ):
11521152 # Pass the wrapper directly; GrainCheckpointHandler will call save_state with the step
1153- save_args_composite ["iter" ] = GrainCheckpointSave (item = data_iterator )
1154- elif not isinstance (data_iterator , list ) and isinstance (data_iterator .local_iterator , ElasticIterator ):
1153+ save_args_composite ["iter" ] = GrainCheckpointSave (item = data_iterator ) # pyrefly: ignore[bad-assignment]
1154+ elif not isinstance (data_iterator , list ) and isinstance (data_iterator .local_iterator , ElasticIterator ): # pyrefly: ignore[missing-attribute]
11551155 # ElasticIterator checkpoints a single global scalar shared by all shards.
1156- save_args_composite ["iter" ] = GrainCheckpointSave (item = data_iterator .local_iterator )
1156+ save_args_composite ["iter" ] = GrainCheckpointSave (item = data_iterator .local_iterator ) # pyrefly: ignore[bad-assignment]
11571157 else :
11581158 if not isinstance (data_iterator , list ):
11591159 data_iterator = [data_iterator ]
@@ -1163,8 +1163,8 @@ def save_checkpoint(checkpoint_manager, step, state, config=None, data_iterator=
11631163 process_count_total = process_count_total // config .expansion_factor_real_data
11641164 for i , data_iter in enumerate (data_iterator ):
11651165 process_index = jax .process_index () + i * jax .process_count ()
1166- grain_iters_to_save .append ((data_iter .local_iterator , process_index , process_count_total ))
1167- save_args_composite ["iter" ] = GrainCheckpointSave (item = grain_iters_to_save )
1166+ grain_iters_to_save .append ((data_iter .local_iterator , process_index , process_count_total )) # pyrefly: ignore[missing-attribute]
1167+ save_args_composite ["iter" ] = GrainCheckpointSave (item = grain_iters_to_save ) # pyrefly: ignore[bad-assignment]
11681168
11691169 custom_metadata = {}
11701170 if config :
0 commit comments