Skip to content

Commit 4145a59

Browse files
h-jooGoogle-ML-Automation
authored andcommitted
Automated Code Change
PiperOrigin-RevId: 944747754
1 parent 0cec2ab commit 4145a59

13 files changed

Lines changed: 51 additions & 51 deletions

src/maxtext/models/deepseek.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -303,9 +303,9 @@ def self_attention_with_norm_op(
303303
return hidden_states, intermediate_inputs
304304

305305
def engram_op(self, x, decoder_input_tokens):
306-
normed_x = self.engram_layer_norm(x)
306+
normed_x = self.engram_layer_norm(x) # pyrefly: ignore[not-callable]
307307
hash_ids = self.ngram_hash_mapping(decoder_input_tokens)[self.layer_idx]
308-
return self.engram(normed_x, hash_ids)
308+
return self.engram(normed_x, hash_ids) # pyrefly: ignore[not-callable]
309309

310310

311311
class DeepSeekDenseLayer(DeepSeekGenericLayer):

src/maxtext/models/gpt3.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -330,7 +330,7 @@ def init_kv_caches(self, inputs_kv_shape: tuple[int, ...]):
330330
)
331331

332332
def update_kv_caches(self, key, value, decoder_segment_ids, model_mode, previous_chunk):
333-
prefill_kv_cache, ar_kv_cache = self.KVCache_0(
333+
prefill_kv_cache, ar_kv_cache = self.KVCache_0( # pyrefly: ignore[not-callable]
334334
key=key,
335335
value=value,
336336
decoder_segment_ids=decoder_segment_ids,

src/maxtext/models/llama2.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -218,8 +218,8 @@ def update_cache(cache, val):
218218
return cache.at[layer_idx].set(val)
219219
return cache
220220

221-
stacked_kv_cache = jax.tree_util.tree_map(update_cache, stacked_kv_cache, kv_cache)
222-
return (layer_output, stacked_kv_cache, layer_idx + 1), None
221+
stacked_kv_cache = jax.tree_util.tree_map(update_cache, stacked_kv_cache, kv_cache) # pyrefly: ignore[unbound-name]
222+
return (layer_output, stacked_kv_cache, layer_idx + 1), None # pyrefly: ignore[unbound-name]
223223
elif cfg.scan_layers:
224224
return layer_output, None
225225
else:

src/maxtext/models/llama4.py

Lines changed: 10 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -50,7 +50,7 @@ class Llama4UnfoldConvolution(nnx.Module):
5050
config: Config containing model parameters
5151
"""
5252

53-
def __init__(self, config: Config, *, rngs: nnx.Rngs = None):
53+
def __init__(self, config: Config, *, rngs: nnx.Rngs = None): # pyrefly: ignore[bad-function-definition]
5454
self.config = config
5555
self.rngs = rngs
5656
self.vit_unfold_linear = linears.DenseGeneral(
@@ -123,7 +123,7 @@ class Llama4VisionMLP(nnx.Module):
123123
config: Config containing model parameters
124124
"""
125125

126-
def __init__(self, config: Config, *, rngs: nnx.Rngs = None):
126+
def __init__(self, config: Config, *, rngs: nnx.Rngs = None): # pyrefly: ignore[bad-function-definition]
127127
self.config = config
128128
self.rngs = rngs
129129
self.vit_encoder_layer_mlp_fc1 = linears.DenseGeneral(
@@ -157,7 +157,7 @@ class Llama4VisionMLP2(nnx.Module):
157157
config: Config containing model parameters
158158
"""
159159

160-
def __init__(self, config: Config, *, rngs: nnx.Rngs = None):
160+
def __init__(self, config: Config, *, rngs: nnx.Rngs = None): # pyrefly: ignore[bad-function-definition]
161161
self.config = config
162162
self.rngs = rngs
163163
self.vit_pixel_shuffle_mlp_fc1 = linears.DenseGeneral(
@@ -196,7 +196,7 @@ class Llama4VisionPixelShuffleMLP(nnx.Module):
196196
config: Config containing model parameters
197197
"""
198198

199-
def __init__(self, config: Config, *, rngs: nnx.Rngs = None):
199+
def __init__(self, config: Config, *, rngs: nnx.Rngs = None): # pyrefly: ignore[bad-function-definition]
200200
self.config = config
201201
self.rngs = rngs
202202
self.pixel_shuffle_ratio = self.config.pixel_shuffle_ratio_for_vit
@@ -221,7 +221,7 @@ class Llama4MultiModalProjector(nnx.Module):
221221
config: Config containing model parameters
222222
"""
223223

224-
def __init__(self, config: Config, mesh: Mesh, *, rngs: nnx.Rngs = None):
224+
def __init__(self, config: Config, mesh: Mesh, *, rngs: nnx.Rngs = None): # pyrefly: ignore[bad-function-definition]
225225
self.config = config
226226
self.mesh = mesh
227227
self.rngs = rngs
@@ -515,8 +515,8 @@ def update_cache(cache, val):
515515
return cache.at[layer_idx].set(val)
516516
return cache
517517

518-
stacked_kv_cache = jax.tree_util.tree_map(update_cache, stacked_kv_cache, kv_cache)
519-
return (layer_output, stacked_kv_cache, layer_idx + 1), None
518+
stacked_kv_cache = jax.tree_util.tree_map(update_cache, stacked_kv_cache, kv_cache) # pyrefly: ignore[unbound-name]
519+
return (layer_output, stacked_kv_cache, layer_idx + 1), None # pyrefly: ignore[unbound-name]
520520
elif cfg.scan_layers:
521521
return layer_output, None
522522
else:
@@ -623,7 +623,7 @@ def __call__(
623623
class Llama4VisionEncoderLayer(nnx.Module):
624624
"""Transformer encoder layer for Llama4 vision model."""
625625

626-
def __init__(self, config: Config, mesh: Mesh, *, rngs: nnx.Rngs = None):
626+
def __init__(self, config: Config, mesh: Mesh, *, rngs: nnx.Rngs = None): # pyrefly: ignore[bad-function-definition]
627627
self.config = config
628628
self.mesh = mesh
629629
self.rngs = rngs
@@ -701,7 +701,7 @@ class Llama4VisionEncoder(nnx.Module):
701701
mesh: Mesh, JAX device mesh (used for sharding)
702702
"""
703703

704-
def __init__(self, config: Config, mesh: Mesh, *, rngs: nnx.Rngs = None):
704+
def __init__(self, config: Config, mesh: Mesh, *, rngs: nnx.Rngs = None): # pyrefly: ignore[bad-function-definition]
705705
self.config = config
706706
self.mesh = mesh
707707
self.rngs = rngs
@@ -733,7 +733,7 @@ class Llama4VisionModel(nnx.Module):
733733
mesh: Mesh, JAX device mesh (used for sharding)
734734
"""
735735

736-
def __init__(self, config: Config, mesh: Mesh, *, rngs: nnx.Rngs = None):
736+
def __init__(self, config: Config, mesh: Mesh, *, rngs: nnx.Rngs = None): # pyrefly: ignore[bad-function-definition]
737737
self.config = config
738738
self.mesh = mesh
739739
self.rngs = rngs

src/maxtext/models/mistral.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -192,8 +192,8 @@ def update_cache(cache, val):
192192
return cache.at[layer_idx].set(val)
193193
return cache
194194

195-
stacked_kv_cache = jax.tree_util.tree_map(update_cache, stacked_kv_cache, kv_cache)
196-
return (layer_output, stacked_kv_cache, layer_idx + 1), None
195+
stacked_kv_cache = jax.tree_util.tree_map(update_cache, stacked_kv_cache, kv_cache) # pyrefly: ignore[unbound-name]
196+
return (layer_output, stacked_kv_cache, layer_idx + 1), None # pyrefly: ignore[unbound-name]
197197
elif cfg.scan_layers:
198198
return layer_output, None
199199
else:

src/maxtext/models/models.py

Lines changed: 9 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -102,9 +102,9 @@ def setup(self):
102102
self.mtp_block = multi_token_prediction_block_as_linen(
103103
config=self.config,
104104
mesh=self.mesh,
105-
transformer_layer_module=mtp_layer_nnx,
106-
decoder=self.decoder,
107-
rngs=self.make_rng("mtp_block"),
105+
transformer_layer_module=mtp_layer_nnx, # pyrefly: ignore[bad-argument-type]
106+
decoder=self.decoder, # pyrefly: ignore[bad-argument-type]
107+
rngs=self.make_rng("mtp_block"), # pyrefly: ignore[bad-argument-type]
108108
)
109109

110110
def logits_from_hidden_states_for_vocab_tiling(self, hidden_states, deterministic, model_mode):
@@ -163,15 +163,15 @@ def __call__(
163163
deepstack_visual_embeds = None
164164

165165
if self.config.use_multimodal and encoder_images is not None:
166-
image_embeddings, deepstack_visual_embeds = self.vision_encoder(
166+
image_embeddings, deepstack_visual_embeds = self.vision_encoder( # pyrefly: ignore[not-callable]
167167
input_images=encoder_images, deterministic=not enable_dropout
168168
)
169169
bidirectional_mask_image = mm_processor.get_bidirectional_mask_vision(
170170
self.config, decoder_input_tokens, is_video=False
171171
)
172172

173173
if self.config.use_multimodal and encoder_videos is not None:
174-
video_embeddings, deepstack_visual_embeds = self.vision_encoder(
174+
video_embeddings, deepstack_visual_embeds = self.vision_encoder( # pyrefly: ignore[not-callable]
175175
input_images=encoder_videos, deterministic=not enable_dropout
176176
)
177177
bidirectional_mask_video = mm_processor.get_bidirectional_mask_vision(
@@ -384,7 +384,7 @@ def __init__(
384384
dummy_attention_metadata = None
385385

386386
if not cfg.pure_nnx_decoder:
387-
self.decoder.lazy_init(
387+
self.decoder.lazy_init( # pyrefly: ignore[missing-attribute]
388388
shared_embedding=self.token_embedder,
389389
decoder_input_tokens=dummy_decoder_input_tokens,
390390
decoder_positions=dummy_decoder_positions,
@@ -490,15 +490,15 @@ def __call__(
490490
audio_embeddings = None
491491
deepstack_visual_embeds = None
492492
if self.config.use_multimodal and encoder_images is not None:
493-
image_embeddings, deepstack_visual_embeds = self.vision_encoder(
493+
image_embeddings, deepstack_visual_embeds = self.vision_encoder( # pyrefly: ignore[not-callable]
494494
input_images=encoder_images, deterministic=not enable_dropout
495495
)
496496
bidirectional_mask_image = mm_processor.get_bidirectional_mask_vision(
497497
self.config, decoder_input_tokens, is_video=False
498498
)
499499

500500
if self.config.use_multimodal and encoder_videos is not None:
501-
video_embeddings, deepstack_visual_embeds = self.vision_encoder(
501+
video_embeddings, deepstack_visual_embeds = self.vision_encoder( # pyrefly: ignore[not-callable]
502502
input_images=encoder_videos, deterministic=not enable_dropout
503503
)
504504
bidirectional_mask_video = mm_processor.get_bidirectional_mask_vision(
@@ -563,7 +563,7 @@ def __call__(
563563
kv_caches=kv_caches,
564564
attention_metadata=attention_metadata,
565565
deepstack_visual_embeds=deepstack_visual_embeds,
566-
mutable=mutable_collections,
566+
mutable=mutable_collections, # pyrefly: ignore[unexpected-keyword]
567567
) # pytype: disable=wrong-keyword-args
568568

569569
# If we are initializing the model AND MTP is enabled, we must create

src/maxtext/utils/gradient_accumulation.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -157,7 +157,7 @@ def reshape_to_microbatch_accumulations(batch_arr):
157157
"ga_params": ga_params,
158158
}
159159
if is_nnx:
160-
init_grad_and_loss["rest_state"] = rest
160+
init_grad_and_loss["rest_state"] = rest # pyrefly: ignore[unbound-name]
161161

162162
grad_and_loss, aux = jax.lax.scan(
163163
accumulate_gradient, init_grad_and_loss, data, length=config.gradient_accumulation_steps

src/maxtext/utils/max_utils.py

Lines changed: 11 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -336,8 +336,8 @@ def initialize_jax_for_gpu(raw_keys):
336336

337337
jax.distributed.initialize(
338338
coordinator_address=f"{coordinator_ip}:{coordinator_port}",
339-
num_processes=int(os.getenv("NNODES")),
340-
process_id=int(os.getenv("NODE_RANK")),
339+
num_processes=int(os.getenv("NNODES")), # pyrefly: ignore[bad-argument-type]
340+
process_id=int(os.getenv("NODE_RANK")), # pyrefly: ignore[bad-argument-type]
341341
initialization_timeout=raw_keys["jax_distributed_initialization_timeout"],
342342
local_device_ids=devices,
343343
)
@@ -349,16 +349,16 @@ def initialize_jax_for_cpu(raw_keys):
349349
coordinator_ip_address = get_coordinator_ip_address()
350350
coordinator_address = coordinator_ip_address + ":1234" # JAX coordinator port used in XPK
351351
# Env variables to be set in XPK or otherwise
352-
job_index = int(os.environ.get("JOB_INDEX"))
353-
job_completion_index = int(os.environ.get("JOB_COMPLETION_INDEX"))
354-
processes_in_job = int(os.environ.get("PROCESSES_IN_JOB"))
352+
job_index = int(os.environ.get("JOB_INDEX")) # pyrefly: ignore[bad-argument-type]
353+
job_completion_index = int(os.environ.get("JOB_COMPLETION_INDEX")) # pyrefly: ignore[bad-argument-type]
354+
processes_in_job = int(os.environ.get("PROCESSES_IN_JOB")) # pyrefly: ignore[bad-argument-type]
355355
pid = job_index * processes_in_job + job_completion_index
356356
max_logging.log(f" Jax process id is {pid} ")
357357
# Explicit initialize is needed only for CPUs
358358
jax.distributed.initialize(
359359
coordinator_address=coordinator_address,
360360
process_id=pid,
361-
num_processes=int(os.environ.get("JAX_PROCESS_COUNT")),
361+
num_processes=int(os.environ.get("JAX_PROCESS_COUNT")), # pyrefly: ignore[bad-argument-type]
362362
initialization_timeout=raw_keys["jax_distributed_initialization_timeout"],
363363
)
364364

@@ -444,7 +444,7 @@ def get_coordinator_ip_address():
444444
max_coordinator_lookups = 50
445445
while not coordinator_found and lookup_attempt <= max_coordinator_lookups:
446446
try:
447-
coordinator_ip_address = socket.gethostbyname(coordinator_address)
447+
coordinator_ip_address = socket.gethostbyname(coordinator_address) # pyrefly: ignore[bad-argument-type]
448448
coordinator_found = True
449449
except socket.gaierror:
450450
max_logging.log(
@@ -676,7 +676,7 @@ def _cross_entropy_with_logits_fwd(logits: jnp.ndarray, targets: jnp.ndarray, z_
676676
log_z = jnp.squeeze(jnp.log(sum_exp) + max_logit, axis=-1)
677677
total_z_loss = z_loss * jax.lax.square(log_z)
678678
loss += total_z_loss
679-
return (loss, total_z_loss), (
679+
return (loss, total_z_loss), ( # pyrefly: ignore[bad-return]
680680
logits,
681681
targets,
682682
z_loss,
@@ -698,11 +698,11 @@ def _cross_entropy_with_logits_bwd(
698698
g: tuple[jnp.ndarray, jnp.ndarray],
699699
) -> tuple[jnp.ndarray, None, None]:
700700
"""Backward-mode of `cross_entropy_with_logits`."""
701-
g = g[0] # Ignore z_loss component as that is only used for logging.
701+
g = g[0] # Ignore z_loss component as that is only used for logging. # pyrefly: ignore[bad-assignment]
702702
logits, targets, z_loss, exp_shifted, sum_exp, log_z = res
703703
# z-loss term adds the (2 * z_loss * log_z) factor.
704704
deriv = jnp.expand_dims(1 + 2 * z_loss * log_z, -1) * exp_shifted / sum_exp - targets
705-
g_logits = jnp.expand_dims(g, axis=-1) * deriv
705+
g_logits = jnp.expand_dims(g, axis=-1) * deriv # pyrefly: ignore[bad-argument-type]
706706

707707
return (
708708
jnp.asarray(g_logits, logits.dtype),
@@ -1230,7 +1230,7 @@ def transformer_engine_context():
12301230
tp_resource="tensor",
12311231
tpsp_resource="tensor_sequence",
12321232
fsdp_resource="fsdp",
1233-
pp_resource=None,
1233+
pp_resource=None, # pyrefly: ignore[bad-argument-type]
12341234
cp_resource="context",
12351235
)
12361236
with global_shard_guard(mesh_resource):

src/maxtext/utils/maxtext_utils.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -95,7 +95,7 @@ def get_functional_train_with_signature(
9595
):
9696
"""Get the shardings (both state and data) for `train_step`."""
9797
functional_train = functools.partial(train_step, model, config, state_mesh_shardings, params_shardings)
98-
functional_train.__name__ = "train_step"
98+
functional_train.__name__ = "train_step" # pyrefly: ignore[missing-attribute]
9999
if config.pure_nnx:
100100
in_shardings = (state_mesh_shardings, data_sharding) # State, batch
101101
else:
@@ -109,7 +109,7 @@ def get_functional_train_with_signature(
109109
def get_functional_eval_with_signature(eval_step, data_sharding, state_mesh_shardings, model, config):
110110
"""Get the shardings (both state and data) for `eval_step`."""
111111
functional_eval = functools.partial(eval_step, model, config)
112-
functional_eval.__name__ = "eval_step"
112+
functional_eval.__name__ = "eval_step" # pyrefly: ignore[missing-attribute]
113113
if config.pure_nnx:
114114
in_shardings = (state_mesh_shardings, data_sharding) # State, batch (NNX: no rng)
115115
else:
@@ -1392,7 +1392,7 @@ def get_abstract_param(model, config):
13921392
{"params": key, "dropout": key, "aqt": key},
13931393
np.ones(input_shape, dtype=jnp.int32),
13941394
np.ones(input_shape, dtype=jnp.int32),
1395-
encoder_images=np.ones(image_shape, dtype=jnp.int32) if config.use_multimodal else None,
1395+
encoder_images=np.ones(image_shape, dtype=jnp.int32) if config.use_multimodal else None, # pyrefly: ignore[no-matching-overload]
13961396
encoder_audios=np.ones(audio_shape, dtype=jnp.float32) if config.use_audio else None,
13971397
)
13981398
return abstract_vars
@@ -2022,7 +2022,7 @@ def to_abstract(x):
20222022
# Convert all input arguments recursively to purely local abstract ShapeDtypeStruct objects
20232023
# to completely bypass remote Array objects and proxy tracing overhead.
20242024
abstract_inputs = jax.tree.map(to_abstract, train_step_inputs)
2025-
p_train_jaxpr = jax.make_jaxpr(unwrapped_step)(*abstract_inputs)
2025+
p_train_jaxpr = jax.make_jaxpr(unwrapped_step)(*abstract_inputs) # pyrefly: ignore[no-matching-overload]
20262026

20272027
local_filename = "train_step.jaxpr"
20282028
local_path = os.path.join(config.dump_jaxpr_local_dir, local_filename)

src/maxtext/utils/maxtext_utils_nnx.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -164,7 +164,7 @@ def create_nnx_sharded_model(
164164
named_sharding = nnx_extract_named_sharding(abstract_state)
165165

166166
if mesh is None:
167-
mesh = abstract_model.mesh
167+
mesh = abstract_model.mesh # pyrefly: ignore[missing-attribute]
168168

169169
# JIT a function that creates the model state with proper sharding from the start.
170170
# By providing out_shardings, we instruct JAX to produce sharded output directly,

0 commit comments

Comments
 (0)