Skip to content

Commit 1008d1e

Browse files
h-jooGoogle-ML-Automation
authored andcommitted
Adding type suppressions for pyrefly
PiperOrigin-RevId: 947186144
1 parent 2707faf commit 1008d1e

2 files changed

Lines changed: 14 additions & 14 deletions

File tree

src/maxtext/common/checkpointing.py

Lines changed: 8 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -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:

src/maxtext/utils/model_creation_utils.py

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -777,7 +777,7 @@ def create_models_and_meshes(trainer_config, sampler_config, trainer_devices, sa
777777
use_no_op_mappings=use_no_op_mappings,
778778
pad_id=tokenizer_pad_id,
779779
)
780-
actor_model.config = None
780+
actor_model.config = None # pyrefly: ignore[missing-attribute]
781781
actor_mesh = reference_mesh
782782
else:
783783
max_logging.log("Creating policy model with same config as reference model on trainer mesh")
@@ -902,7 +902,7 @@ def from_pretrained(
902902
load_parameters_path = epath.Path(config.base_output_directory) / "0" / "items"
903903
# Create a copied Pydantic model with the updated values
904904
pydantic_config = getattr(config, "_pydantic_config", config)
905-
new_config = pydantic_config.model_copy(
905+
new_config = pydantic_config.model_copy( # pyrefly: ignore[missing-attribute]
906906
update={
907907
"load_parameters_path": load_parameters_path,
908908
}
@@ -930,7 +930,7 @@ def from_pretrained(
930930
sharded_state = nnx.state(model)
931931

932932
if mesh is None:
933-
mesh = model.mesh
933+
mesh = model.mesh # pyrefly: ignore[missing-attribute]
934934

935935
with mesh:
936936
if config.load_parameters_path:
@@ -1114,7 +1114,7 @@ def _free_device_memory(node):
11141114
)
11151115

11161116
if is_nnx_checkpoint:
1117-
restored_root = restored["base"] if has_base_key else restored
1117+
restored_root = restored["base"] if has_base_key else restored # pyrefly: ignore[unbound-name]
11181118
checkpoint = jax.tree.map(
11191119
lambda v: v["value"],
11201120
restored_root,
@@ -1202,11 +1202,11 @@ def _walk_align(ckpt, model_arr, axes):
12021202
with mesh:
12031203
use_no_op_mappings = "maxtext_config" in config.vllm_additional_config
12041204
model = TunixMaxTextAdapter(
1205-
base_model=model,
1205+
base_model=model, # pyrefly: ignore[bad-argument-type]
12061206
use_no_op_mappings=use_no_op_mappings,
12071207
pad_id=tokenizer_pad_id,
12081208
)
1209-
model.config = None
1209+
model.config = None # pyrefly: ignore[missing-attribute]
12101210

12111211
if original_mesh:
12121212
return model

0 commit comments

Comments
 (0)