AI-Hypercomputer / AI-Hypercomputer/maxtext

Memory Leak when initialising Qwen model training with multihost runner

Aberta
#2,374 1 comentário 0 reações 1 responsável Reivindicada por @parambole Ver no GitHub
bug
Linguagem predominante
Python
Estrelas
2.4k
Forks
607
Merge médio
2d 19h
PRs com merge (30d)
158

Descrição

### Bug report

Problems occur for both Qwen Dense (4B) and Qwen MoE (30B-A3B).

Step 1: Convert the Dense and MoE models according to the end-to-end tutorial

Step 2: Initialise a Spot VM and install MaxText

Step 3: Use multihost_runner to initialise a training on a v6e-32 TPU VM from my local machine

For the MoE model
```bash
python3 multihost_runner.py \
--USE_EXISTING_FOLDER="True" \
--RUN_NAME=$RUN_NAME \
--TPU_PREFIX=$TPU_PREFIX \
--COMMAND="source ~/maxtext_venv/bin/activate && python3 -m MaxText.train \
src/MaxText/configs/base.yml \
run_name="qwen3_30B_A3B-finetune-maxtext" \
base_output_directory=gs://${my_bucket}/qwen3_maxtext_ckpt/ \
load_parameters_path=gs://${my_bucket}/qwen3_maxtext_ckpt/Qwen3-30B-A3B-Base/0/items \
model_name='qwen3-30b-a3b' \
async_checkpointing=False \
tokenizer_type=huggingface \
tokenizer_path=src/MaxText/assets/qwen3-tokenizer \
dataset_type=grain \
grain_file_type=parquet\
grain_train_files=/tmp/gcsfuse/data_test/part*.parquet \
grain_worker_count=2 \
steps=10 \
sparse_matmul=False \
scan_layers=False \
per_device_batch_size=2 \
max_target_length=2048 \
megablox=False \
ici_fsdp_parallelism=32
"
```
For the Dense model
```bash
python3 multihost_runner.py \
--USE_EXISTING_FOLDER="True" \
--RUN_NAME=$RUN_NAME \
--TPU_PREFIX=$TPU_PREFIX \
--COMMAND="source ~/maxtext_venv/bin/activate && python3 -m MaxText.train \
src/MaxText/configs/base.yml \
run_name="qwen3_4B-finetune-maxtext" \
base_output_directory=gs://${my_bucket}/qwen3_maxtext_ckpt/ \
load_parameters_path=gs://${my_bucket}/qwen3_maxtext_ckpt/Qwen3-4B-base/0/items \
model_name='qwen3-4b' \
async_checkpointing=True \
tokenizer_path=src/MaxText/assets/qwen3-tokenizer \
dataset_type=grain \
grain_file_type=parquet\
grain_train_files=/tmp/gcsfuse/data_test/part*.parquet \
grain_worker_count=2 \
steps=10 \
sparse_matmul=False \
scan_layers=False \
per_device_batch_size=6 \
max_target_length=4096 \
megablox=False \
ici_fsdp_parallelism=32 \
"
```

The dense command works perfectly in a single host VM environment.
But then I used the multihost runner to execute the command.

The TPU VM initialised the model, but then got stuck for more than 10 minutes.

I sshed into the master node and used htop to inspect what was happening

The worker is using 100% single CPU. Then used up all the 700GB of memory and crashed.

### Logs/Output

Using ssh batch size of 1. Attempting to SSH into 1 nodes with a total of 1 workers.
SSH: Attempting to connect to worker 0...
Searching for existing processes on device vfio/...
No existing processes found, so your TPU is ready to use!
2025-09-20 10:00:12.406166: E external/local_xla/xla/stream_executor/cuda/cuda_fft.cc:467] Unable to register cuFFT factory: Attempting to register factory for plugin cuFFT when one has already been registered
WARNING: All log messages before absl::InitializeLog() is called are written to STDERR
E0000 00:00:1758362412.417293 17525 cuda_dnn.cc:8579] Unable to register cuDNN factory: Attempting to register factory for plugin cuDNN when one has already been registered
E0000 00:00:1758362412.420667 17525 cuda_blas.cc:1407] Unable to register cuBLAS factory: Attempting to register factory for plugin cuBLAS when one has already been registered
W0000 00:00:1758362412.431158 17525 computation_placer.cc:177] computation placer already registered. Please check linkage and avoid linking the same target more than once.
W0000 00:00:1758362412.431170 17525 computation_placer.cc:177] computation placer already registered. Please check linkage and avoid linking the same target more than once.
W0000 00:00:1758362412.431172 17525 computation_placer.cc:177] computation placer already registered. Please check linkage and avoid linking the same target more than once.
W0000 00:00:1758362412.431174 17525 computation_placer.cc:177] computation placer already registered. Please check linkage and avoid linking the same target more than once.
2025-09-20 10:00:20.175093: E external/local_xla/xla/stream_executor/cuda/cuda_platform.cc:51] failed call to cuInit: INTERNAL: CUDA error: Failed call to cuInit: UNKNOWN ERROR (303)
Updating keys from env and command line: ['run_name', 'model_name', 'load_parameters_path', 'async_checkpointing', 'megablox', 'sparse_matmul', 'scan_layers', 'base_output_directory', 'ici_fsdp_parallelism', 'tokenizer_path', 'tokenizer_type', 'per_device_batch_size', 'dataset_type', 'grain_train_files', 'grain_file_type', 'grain_worker_count', 'steps', 'max_target_length']
Running Model: qwen3-30b-a3b
Updating following parameters in config

decoder_block: qwen3_moe
base_emb_dim: 2048
base_mlp_dim: 768
base_num_query_heads: 32
base_num_kv_heads: 4
base_num_decoder_layers: 48
head_dim: 128
mlp_activations: ['silu', 'linear']
vocab_size: 151936
normalization_layer_epsilon: 1e-06
use_qk_norm: True
num_experts: 128
num_experts_per_tok: 8
base_moe_mlp_dim: 768
norm_topk_prob: True
rope_max_timescale: 10000000
enable_dropout: False
Updating keys from model: ['decoder_block', 'base_emb_dim', 'base_mlp_dim', 'base_num_query_heads', 'base_num_kv_heads', 'base_num_decoder_layers', 'head_dim', 'mlp_activations', 'vocab_size', 'normalization_layer_epsilon', 'use_qk_norm', 'num_experts', 'num_experts_per_tok', 'base_moe_mlp_dim', 'norm_topk_prob', 'rope_max_timescale', 'enable_dropout']
Attempting to initialize the jax distributed system...
INFO:2025-09-20 10:00:26,474:jax._src.distributed:145: Starting JAX distributed service on [::]:8476
I0920 10:00:26.474668 139574753909888 distributed.py:145] Starting JAX distributed service on [::]:8476
INFO:2025-09-20 10:00:26,476:jax._src.distributed:161: Connecting to JAX distributed service on 10.164.15.194:8476
I0920 10:00:26.476343 139574753909888 distributed.py:161] Connecting to JAX distributed service on 10.164.15.194:8476
Jax distributed system initialized!
Not using emergency checkpoint, ignoring local_checkpoint_directory, local_checkpoint_period, use_replicator_service and replicator_backup_interval_minutes
dataset_type set to grain, will use keys['grain_train_files']='/tmp/gcsfuse/data_test/part*.parquet' as data files, and 2 workers
Config param activations_in_float32: False
Config param adam_b1: 0.9
Config param adam_b2: 0.95
Config param adam_eps: 1e-08
Config param adam_eps_root: 0.0
Config param adam_weight_decay: 0.1
Config param add_bos: True
Config param add_eos: True
Config param allow_split_physical_axes: False
Config param ar_cache_axis_order: 1,2,0,3
Config param async_checkpointing: False
Config param attention: autoselected
Config param attention_bias: False
Config param attention_sink: False
Config param attention_type: global
Config param attn_logits_soft_cap: None
Config param autoregressive_decode_assert:
Config param base_emb_dim: 2048
Config param base_mlp_dim: 768
Config param base_moe_mlp_dim: 768
Config param base_num_decoder_layers: 48
Config param base_num_kv_heads: 4
Config param base_num_query_heads: 32
Config param base_output_directory: gs://cantonese_llm/qwen3_maxtext_ckpt/
Config param beta_fast: 32
Config param beta_slow: 1
Config param capacity_factor: -1.0
Config param cast_logits_to_fp32: True
Config param checkpoint_conversion_fn: None
Config param checkpoint_dir: gs://cantonese_llm/qwen3_maxtext_ckpt/qwen3_30B_A3B-finetune-maxtext/checkpoints/
Config param checkpoint_is_quantized: False
Config param checkpoint_period: 10000
Config param checkpoint_storage_concurrent_gb: 96
Config param checkpoint_storage_target_data_file_size_bytes: 2147483648
Config param checkpoint_storage_use_ocdbt: True
Config param checkpoint_storage_use_zarr3: True
Config param chunk_attn_window_size: 0
Config param collect_stack_trace: False
Config param colocated_python_data_input: False
Config param compile_topology:
Config param compile_topology_num_slices: -1
Config param compiled_trainstep_file:
Config param compute_axis_order: 0,1,2,3
Config param constant_bound_config: []
Config param context: remat
Config param context_parallel_load_balance: True
Config param context_parallel_size: 1
Config param conv_stride_for_vit: 14
Config param cosine_learning_rate_final_fraction: 0.1
Config param custom_mesh:
Config param data_sharding: (('data', 'stage', 'fsdp', 'fsdp_transpose', 'sequence', 'context', 'context_autoregressive', 'tensor', 'tensor_transpose', 'tensor_sequence', 'expert', 'autoregressive'),)
Config param data_shuffle_seed: 0
Config param dataset_name: c4/en:3.0.1
Config param dataset_path:
Config param dataset_type: grain
Config param dcn_autoregressive_parallelism: 1
Config param dcn_context_autoregressive_parallelism: 1
Config param dcn_context_parallelism: 1
Config param dcn_data_parallelism: -1
Config param dcn_expert_parallelism: 1
Config param dcn_fsdp_parallelism: 1
Config param dcn_fsdp_transpose_parallelism: 1
Config param dcn_parallelism: [-1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1]
Config param dcn_pipeline_parallelism: 1
Config param dcn_sequence_parallelism: 1
Config param dcn_tensor_parallelism: 1
Config param dcn_tensor_sequence_parallelism: 1
Config param dcn_tensor_transpose_parallelism: 1
Config param decode_sampling_nucleus_p: -1
Config param decode_sampling_strategy: greedy
Config param decode_sampling_temperature: 1.0
Config param decode_sampling_top_k: 0
Config param decoder_block: DecoderBlockType.QWEN3_MOE
Config param decoder_layer_input: device
Config param dpo_beta: 0.1
Config param dpo_label_smoothing: 0.0
Config param dropout_rate: 0.0
Config param dtype: bfloat16
Config param dtype_mm: float32
Config param dump_hlo: False
Config param dump_hlo_delete_local_after: True
Config param dump_hlo_gcs_dir:
Config param dump_hlo_local_dir: /tmp/xla_dump/
Config param dump_hlo_local_module_name: jit_train_step
Config param dump_hlo_module_name: jit_train_step
Config param dump_hlo_upload_all: False
Config param dump_hlo_xla_flags:
Config param dump_step: -1
Config param emb_dim: 2048
Config param enable_checkpoint_cloud_logger: False
Config param enable_checkpointing: True
Config param enable_data_shuffling: True
Config param enable_dropout: False
Config param enable_emergency_checkpoint: False
Config param enable_gcp_goodput_metrics: True
Config param enable_gcp_step_deviation_metrics: True
Config param enable_goodput_recording: False
Config param enable_jax_profiler: False
Config param enable_llm_inference_pool: False
Config param enable_model_warmup: False
Config param enable_multi_tier_checkpointing: False
Config param enable_nnx: False
Config param enable_orbax_v1: False
Config param enable_padding_causal_mask: True
Config param enable_pathways_goodput: False
Config param enable_prefix_caching: False
Config param enable_single_controller: False
Config param enable_single_replica_ckpt_restoring: False
Config param enable_tensorboard: True
Config param eval_data_columns: ['text']
Config param eval_dataset_name: c4/en:3.0.1
Config param eval_image_column: image
Config param eval_interval: -1
Config param eval_per_device_batch_size: 2.0
Config param eval_split: validation
Config param eval_steps: -1
Config param expansion_factor_real_data: -1
Config param expert_shard_attention_option: fsdp
Config param final_logits_soft_cap: None
Config param first_num_dense_layers: 0
Config param float32_logits: False
Config param float32_qk_product: False
Config param force_unroll: False
Config param freeze_vision_encoder_params: True
Config param fused_mlp: False
Config param fused_qkv: False
Config param gcs_metrics: False
Config param generate_padding_batch_eval: False
Config param generate_padding_batch_train: False
Config param generate_slice: v5e-16
Config param global_batch_size_to_eval_on: 64
Config param global_batch_size_to_load: 64
Config param global_batch_size_to_load_eval: 64
Config param global_batch_size_to_train_on: 64
Config param global_parameter_scale: 1
Config param goodput_upload_interval_seconds: 30
Config param gradient_accumulation_steps: 1
Config param gradient_clipping_threshold: 1.0
Config param grain_eval_files:
Config param grain_file_type: parquet
Config param grain_train_files: /tmp/gcsfuse/data_test/part*.parquet
Config param grain_worker_count: 2
Config param grain_worker_count_eval: 1
Config param hardware: tpu
Config param head_dim: 128
Config param heartbeat_reporting_interval_in_seconds: 5
Config param hf_data_dir:
Config param hf_eval_files:
Config param hf_eval_split:
Config param hf_path:
Config param hf_train_files:
Config param hidden_size_for_vit: 1408
Config param ici_autoregressive_parallelism: 1
Config param ici_context_autoregressive_parallelism: 1
Config param ici_context_parallelism: 1
Config param ici_data_parallelism: 1
Config param ici_expert_parallelism: 1
Config param ici_fsdp_parallelism: 32
Config param ici_fsdp_transpose_parallelism: 1
Config param ici_parallelism: [1, 1, 32, 1, 1, 1, 1, 1, 1, 1, 1, 1]
Config param ici_pipeline_parallelism: 1
Config param ici_sequence_parallelism: 1
Config param ici_tensor_parallelism: 1
Config param ici_tensor_sequence_parallelism: 1
Config param ici_tensor_transpose_parallelism: 1
Config param image_path:
Config param image_placeholder: <|image|>
Config param image_size_for_vit: 896
Config param inference_benchmark_test: False
Config param inference_metadata_file:
Config param inference_microbenchmark_log_file_path:
Config param inference_microbenchmark_loop_iters: 10
Config param inference_microbenchmark_num_samples: [1, 2, 3, 4, 5]
Config param inference_microbenchmark_prefill_lengths: 64,128,256,512,1024
Config param inference_microbenchmark_stages: prefill,generate
Config param inference_server: MaxtextInterleavedServer
Config param inhomogeneous_layer_cycle_interval: 1
Config param init_weights_seed: 0
Config param input_data_sharding_logical_axes: ['activation_embed_and_logits_batch', 'activation_norm_length']
Config param interleave_moe_layer_step: 1
Config param intermediate_size_for_vit: 5632
Config param jax_cache_dir: ~/jax_cache
Config param jax_debug_log_modules:
Config param jax_distributed_initialization_timeout: 300
Config param jax_profiler_port: 9999
Config param key_proj: remat
Config param kv_lora_rank: 512
Config param kv_quant_axis: heads_and_dkv
Config param kv_quant_dtype: int8
Config param learning_rate: 3e-05
Config param learning_rate_schedule_steps: 10
Config param load_balance_loss_weight: 0.01
Config param load_from_prefill_dir: False
Config param load_full_state_path:
Config param load_parameters_path: gs://cantonese_llm/qwen3_maxtext_ckpt/Qwen3-30B-A3B-Base/0/items
Config param local_checkpoint_directory:
Config param local_checkpoint_period: 0
Config param local_rope_max_timescale: -1
Config param log_config: True
Config param log_period: 100
Config param logical_axis_rules: (('activation_batch', ('data', 'fsdp', 'fsdp_transpose', 'expert')), ('activation_batch_no_exp', ('data', 'fsdp', 'fsdp_transpose')), ('activation_embed_and_logits_batch', ('data', 'stage', 'fsdp', 'fsdp_transpose', 'expert')), ('activation_heads', ('tensor', 'tensor_transpose', 'sequence', 'tensor_sequence', 'autoregressive')), ('activation_kv_heads', ('tensor', 'tensor_transpose', 'sequence', 'tensor_sequence')), ('activation_length', ('sequence', 'context', 'expert')), ('activation_length', ('context', 'expert')), ('activation_length_no_exp', ('sequence', 'context')), ('activation_length_no_exp', ('context',)), ('activation_norm_length', ('tensor_sequence', 'context', 'sequence')), ('activation_q_length', ('context', 'expert')), ('activation_q_length_no_exp', ('context',)), ('prefill_activation_length', ('sequence', 'context')), ('prefill_activation_norm_length', ('tensor_sequence', 'context', 'sequence')), ('activation_kv_length', ()), ('activation_embed', ('tensor', 'tensor_transpose')), ('activation_mlp', ('tensor', 'tensor_transpose', 'tensor_sequence')), ('activation_kv', ('tensor', 'tensor_transpose', 'tensor_sequence')), ('activation_prefill_kv_batch', ('data', 'fsdp', 'fsdp_transpose', 'expert')), ('activation_kv_batch', ('data', 'fsdp', 'fsdp_transpose', 'expert')), ('activation_kv_batch_no_exp', ('data', 'fsdp', 'fsdp_transpose')), ('activation_kv_head_dim', ('tensor', 'tensor_transpose', 'tensor_sequence')), ('activation_vocab', ('tensor', 'tensor_transpose', 'sequence', 'tensor_sequence')), ('activation_vocab', ('tensor', 'tensor_transpose')), ('activation_vocab', 'tensor_sequence'), ('activation_vocab', ('sequence', 'context')), ('activation_stage', 'stage'), ('activation_exp', ('expert',)), ('decode_batch', ('data', 'fsdp', 'fsdp_transpose', 'expert')), ('decode_length', ('sequence',)), ('mlp', ('fsdp_transpose', 'tensor', 'tensor_sequence', 'autoregressive')), ('mlp_no_fsdp', ('tensor', 'tensor_sequence', 'autoregressive')), ('vocab', ('tensor', 'tensor_transpose', 'tensor_sequence', 'autoregressive')), ('heads', ('tensor', 'tensor_transpose', 'tensor_sequence', 'autoregressive')), ('q_heads', ('tensor', 'tensor_transpose', 'tensor_sequence', 'autoregressive')), ('kv_heads', ('tensor', 'tensor_transpose', 'tensor_sequence', 'autoregressive')), ('embed', ('fsdp', 'fsdp_transpose', 'sequence', 'tensor_transpose', 'context', 'expert')), ('embed', ('fsdp', 'sequence', 'tensor_transpose', 'context', 'expert')), ('embed', ('fsdp', 'fsdp_transpose', 'sequence', 'context', 'expert')), ('embed', ('fsdp', 'sequence', 'context', 'expert')), ('embed_no_exp', ('fsdp', 'fsdp_transpose', 'sequence', 'tensor_transpose', 'context')), ('embed_no_exp', ('fsdp', 'sequence', 'tensor_transpose', 'context')), ('embed_no_exp', ('fsdp', 'fsdp_transpose', 'sequence', 'context')), ('embed_no_exp', ('fsdp', 'sequence', 'context')), ('embed_tensor_transpose', ('tensor_transpose',)), ('q_lora', ('fsdp', 'fsdp_transpose', 'sequence', 'context', 'tensor_transpose', 'expert')), ('q_lora', ('fsdp', 'sequence', 'context', 'tensor_transpose', 'expert')), ('q_lora', ('fsdp', 'fsdp_transpose', 'sequence', 'context', 'expert')), ('q_lora', ('fsdp', 'sequence', 'context', 'expert')), ('kv_lora', ('fsdp', 'fsdp_transpose', 'sequence', 'context', 'tensor_transpose', 'expert')), ('kv_lora', ('fsdp', 'sequence', 'context', 'tensor_transpose', 'expert')), ('kv_lora', ('fsdp', 'fsdp_transpose', 'sequence', 'context', 'expert')), ('kv_lora', ('fsdp', 'sequence', 'context', 'expert')), ('norm', ('tensor', 'tensor_transpose')), ('layers', 'stage'), ('kv', ()), ('kv_head_dim', ()), ('cache_batch_prefill', ()), ('cache_batch', ()), ('cache_heads_none', ()), ('cache_heads', ('autoregressive', 'tensor', 'tensor_transpose', 'tensor_sequence')), ('cache_heads', ('autoregressive', 'tensor', 'tensor_sequence')), ('cache_kv', ()), ('cache_sequence', ()), ('exp', 'expert'), ('paged_kv_heads', ('tensor',)), ('num_pages', ()), ('tokens_per_page', ()), ('paged_kv_head_dim_size', ()))
Config param logits_dot_in_fp32: False
Config param logits_via_embedding: False
Config param lora_input_adapters_path:
Config param matmul_precision: default
Config param max_checkify: False
Config param max_corpus_chars: 10000000
Config param max_num_images_per_example: -1
Config param max_position_embeddings: 163840
Config param max_prefill_predict_length: 64
Config param max_target_length: 2048
Config param megablox: False
Config param mesh_axes: ['data', 'stage', 'fsdp', 'fsdp_transpose', 'sequence', 'context', 'context_autoregressive', 'tensor', 'tensor_transpose', 'tensor_sequence', 'expert', 'autoregressive']
Config param metrics_dir: gs://cantonese_llm/qwen3_maxtext_ckpt/qwen3_30B_A3B-finetune-maxtext/metrics/
Config param metrics_file:
Config param micro_batch_size_to_eval_on: 64
Config param micro_batch_size_to_train_on: 64
Config param mla_naive_kvcache: True
Config param mlp_activations: ['silu', 'linear']
Config param mlp_activations_limit: -1.0
Config param mlp_bias: False
Config param mlp_dim: 768
Config param mlpwi: remat
Config param mlpwi_0: remat
Config param mlpwi_1: remat
Config param mlpwo: remat
Config param model_call_mode:
Config param model_fsdp_ag_once: False
Config param model_name: qwen3-30b-a3b
Config param moe_fsdp_use_two_stage_all_gather: False
Config param moe_mlp_dim: 768
Config param monitor_goodput: False
Config param monitor_step_time_deviation: True
Config param mscale: 1.0
Config param mtp_eval_target_module: 0
Config param mtp_loss_scaling_factor: 0.1
Config param mtp_num_layers: 0
Config param mu_dtype: float32
Config param multi_sampling: False
Config param multi_tier_checkpointing_backup_interval_minutes: 0
Config param n_routing_groups: -1
Config param nope_layer_interval: -1
Config param norm_topk_prob: True
Config param normalization_layer_epsilon: 1e-06
Config param normalize_embedding_logits: True
Config param num_attention_heads_for_vit: 16
Config param num_channels_for_vit: 3
Config param num_decoder_layers: 48
Config param num_epoch: 1
Config param num_experts: 128
Config param num_experts_per_tok: 8
Config param num_hidden_layers_for_vit: 34
Config param num_kv_heads: 4
Config param num_layers_per_pipeline_stage: 1
Config param num_pipeline_microbatches: -1
Config param num_pipeline_repeats: -1
Config param num_query_heads: 32
Config param num_slices: 1
Config param opt_type: adamw
Config param optimize_mesh_for_tpu_v6e: False
Config param optimizer_memory_host_offload: False
Config param original_max_position_embeddings: 4096
Config param out_proj: remat
Config param override_model_config: False
Config param packing: True
Config param pagedattn_head_dim_alignment: 128
Config param pagedattn_max_pages_per_group: 64
Config param pagedattn_num_pages: 64
Config param pagedattn_pages_per_compute_block: 4
Config param pagedattn_tokens_per_page: 32
Config param param_scan_axis: 1
Config param parameter_memory_host_offload: False
Config param patch_size_for_vit: 14
Config param per_device_batch_size: 2.0
Config param pipeline_delay_activation_forwarding: False
Config param pipeline_fsdp_ag_once: False
Config param pipeline_parallel_layers: -1
Config param pixel_shuffle_ratio_for_vit: 0.5
Config param posemb_type_for_vit: learn
Config param prefill_cache_axis_order: 1,2,0,3
Config param prefill_cache_dir:
Config param prefill_chunk_size: 256
Config param prefill_slice: v5e-16
Config param prefix_caching_dram_byte: 100000000000
Config param prefix_caching_hbm_byte: 10000000000
Config param profile_cleanly: True
Config param profile_periodically_period: -1
Config param profiler:
Config param profiler_steps: 5
Config param projector_dropout_for_vit: 0.0
Config param projector_input_dim_for_vit: 4096
Config param projector_output_dim_for_vit: 4096
Config param prometheus_port: 0
Config param prompt: I love to
Config param q_lora_rank: 0
Config param qk_nope_head_dim: 128
Config param qk_rope_head_dim: 64
Config param qkv_proj: remat
Config param quant_cfg_path:
Config param quantization:
Config param quantization_calibration_method: absmax
Config param quantization_local_shard_count: 1
Config param quantize_kvcache: False
Config param query_proj: remat
Config param ragged_block_size: 256
Config param record_internal_nn_metrics: 0
Config param remat_policy: full
Config param remat_policy_for_vit: minimal
Config param replicate_quant_scale: False
Config param report_heartbeat_metric_for_gcp_monitoring: False
Config param report_performance_metric_for_gcp_monitoring: False
Config param reshape_q: False
Config param return_log_prob: False
Config param reuse_example_batch: 0
Config param rope_attention_scaling: False
Config param rope_factor: 40
Config param rope_interleave: True
Config param rope_linear_scaling_factor: 1.0
Config param rope_max_timescale: 10000000
Config param rope_min_timescale: 1
Config param rope_theta_for_vit: 10000
Config param rope_truncate: True
Config param rope_type: default
Config param rope_use_scale: True
Config param routed_bias: False
Config param routed_scaling_factor: 1.0
Config param routed_score_func:
Config param run_name: qwen3_30B_A3B-finetune-maxtext
Config param sa_block_kv: 512
Config param sa_block_kv_compute: 512
Config param sa_block_kv_dkv: 512
Config param sa_block_kv_dkv_compute: 512
Config param sa_block_kv_dq: 512
Config param sa_block_q: 512
Config param sa_block_q_dkv: 512
Config param sa_block_q_dq: 512
Config param sa_k_layout: HEAD_DIM_MINOR
Config param sa_q_layout: HEAD_DIM_MINOR
Config param sa_use_fused_bwd_kernel: False
Config param sa_v_layout: HEAD_DIM_MINOR
Config param save_checkpoint_on_completion: True
Config param save_config_to_gcs: False
Config param save_quantized_params_path:
Config param scan_layers: False
Config param scan_layers_per_stage: False
Config param scan_pipeline_iterations: True
Config param set_remat_policy_on_layers_per_stage: False
Config param set_remat_policy_on_pipeline_iterations: True
Config param sft_train_on_completion_only: False
Config param sharding_tolerance: 0.02
Config param shardy: True
Config param shared_experts: 1
Config param skip_first_n_steps_for_profiler: 1
Config param skip_jax_distributed_system: False
Config param sliding_window_size: 0
Config param source_checkpoint_layout: orbax
Config param sparse_matmul: False
Config param stack_prefill_result_cache: False
Config param stack_trace_interval_seconds: 600
Config param stack_trace_to_cloud: False
Config param step_deviation_interval_seconds: 30
Config param steps: 10
Config param subslice_shape:
Config param target_eval_loss: 0.0
Config param temperature_tuning: False
Config param tensorboard_dir: gs://cantonese_llm/qwen3_maxtext_ckpt/qwen3_30B_A3B-finetune-maxtext/tensorboard/
Config param tile_activation_dim: 1024
Config param tile_batch_seq: 512
Config param tile_size_for_vit: 336
Config param tile_weight_dim: 1024
Config param tokenize_eval_data: True
Config param tokenize_train_data: True
Config param tokenizer_path: src/MaxText/assets/qwen3-tokenizer
Config param tokenizer_type: huggingface
Config param topk_routing_group: -1
Config param train_data_columns: ['text']
Config param train_image_column: image
Config param train_split: train
Config param trainable_position_size: -1
Config param upload_all_profiler_results: False
Config param use_chat_template: False
Config param use_chunked_prefill: False
Config param use_custom_sort_vjp: False
Config param use_dpo: False
Config param use_iota_embed: False
Config param use_multimodal: False
Config param use_post_attn_norm: False
Config param use_post_ffw_norm: False
Config param use_qk_norm: True
Config param use_qwix_quantization: False
Config param use_ragged_attention: False
Config param use_random_routing: False
Config param use_ring_of_experts: False
Config param use_sft: False
Config param use_untrainable_positional_embedding: False
Config param use_vertex_tensorboard: False
Config param using_pipeline_parallelism: False
Config param v_head_dim: 128
Config param value_proj: remat
Config param vertex_tensorboard_project:
Config param vertex_tensorboard_region:
Config param vision_output_dim_for_vit: 4096
Config param vocab_size: 151936
Config param warmup_steps_fraction: 0.1
Config param weight_dtype: float32
System Information: Jax Version: 0.7.2
System Information: Jaxlib Version: 0.7.2
System Information: Jax Backend: PJRT C API
TFRT TPU v6 lite
Built on Sep 11 2025 15:57:19 (1757631439) cl/804429027
Num_devices: 32, shape (1, 1, 32, 1, 1, 1, 1, 1, 1, 1, 1, 1)
Setting up checkpoint logger...
Creating checkpoint manager with ocdbt=True and zarr3=True
I0920 10:00:35.203486 139574753909888 base_pytree_checkpoint_handler.py:385] Created BasePyTreeCheckpointHandler: use_ocdbt=True, use_zarr3=True, pytree_metadata_options=PyTreeMetadataOptions(support_rich_types=False), array_metadata_store=, enable_pinned_host_transfer=False, save_concurrent_bytes: 96000000000 (89.4 GiB), restore_concurrent_bytes: 96000000000 (89.4 GiB)
I0920 10:00:35.203660 139574753909888 checkpoint_manager.py:694] [process=6][thread=MainThread] CheckpointManager init: checkpointers=None, item_names=('items', 'iter'), item_handlers={'items': }, handler_registry=None
I0920 10:00:35.203894 139574753909888 composite_checkpoint_handler.py:237] Deferred registration for item: "items". Adding handler `` for item "items" and save args `` and restore args `` to `_handler_registry`.
I0920 10:00:35.203937 139574753909888 composite_checkpoint_handler.py:237] Deferred registration for item: "metrics". Adding handler `` for item "metrics" and save args `` and restore args `` to `_handler_registry`.
I0920 10:00:35.203969 139574753909888 composite_checkpoint_handler.py:505] Initialized registry DefaultCheckpointHandlerRegistry({('items', ): , ('items', ): , ('metrics', ): , ('metrics', ): }).
I0920 10:00:35.204277 139574753909888 abstract_checkpointer.py:35] orbax-checkpoint version: 0.11.25
I0920 10:00:35.557745 139574753909888 checkpoint_manager.py:1757] Found 0 checkpoint steps in gs://cantonese_llm/qwen3_maxtext_ckpt/qwen3_30B_A3B-finetune-maxtext/checkpoints
Checkpoint manager created!
I0920 10:00:35.579838 139574753909888 checkpoint_manager.py:911] [process=6][thread=MainThread] CheckpointManager created, primary_host=0, CheckpointManagerOptions=CheckpointManagerOptions(save_interval_steps=10000, max_to_keep=None, keep_time_interval=None, keep_period=None, should_keep_fn=None, best_fn=None, best_mode='max', keep_checkpoints_without_metrics=True, step_prefix=None, step_format_fixed_length=None, step_name_format=None, create=True, cleanup_tmp_directories=False, save_on_steps=frozenset(), single_host_load_and_broadcast=False, todelete_subdir=None, todelete_full_path=None, enable_hns=False, enable_background_delete=False, read_only=False, enable_async_checkpointing=False, async_options=None, multiprocessing_options=MultiprocessingOptions(primary_host=0, active_processes=None, barrier_sync_key_prefix=None), should_save_fn=None, file_options=FileOptions(path_permission_mode=None), save_root_metadata=True, temporary_path_class=None, save_decision_policy=None, preservation_policy=None, prevent_write_metrics=False, enable_should_save_is_saving_in_progress_check=True, enable_per_process_directory_creation=False), root_directory=gs://cantonese_llm/qwen3_maxtext_ckpt/qwen3_30B_A3B-finetune-maxtext/checkpoints:
Found 16 files for train/eval with grain
Tokenizer path: src/MaxText/assets/qwen3-tokenizer
Loading HF tokenizer: src/MaxText/assets/qwen3-tokenizer
checkpoint manager exists so trying to load this run's existing checkpoint
restoring params from gs://cantonese_llm/qwen3_maxtext_ckpt/Qwen3-30B-A3B-Base/0/items
Creating checkpoint manager with ocdbt=True and zarr3=True
I0920 10:00:40.036749 139574753909888 base_pytree_checkpoint_handler.py:385] Created BasePyTreeCheckpointHandler: use_ocdbt=True, use_zarr3=True, pytree_metadata_options=PyTreeMetadataOptions(support_rich_types=False), array_metadata_store=, enable_pinned_host_transfer=False, save_concurrent_bytes: 96000000000 (89.4 GiB), restore_concurrent_bytes: 96000000000 (89.4 GiB)
I0920 10:00:40.143391 139574753909888 checkpointer.py:304] Restoring checkpoint from gs://cantonese_llm/qwen3_maxtext_ckpt/Qwen3-30B-A3B-Base/0/items.
I0920 10:00:40.405811 18851 google_auth_provider.cc:181] Running on GCE, using service account tpu-service-account@XXXXX.iam.gserviceaccount.com
W0920 10:00:42.614231 139574753909888 transform_utils.py:230] The transformations API will eventually be replaced by an upgraded design. The current API will not be removed until this point, but it will no longer be actively worked on.
I0920 10:00:42.615916 139574753909888 transform_utils.py:288] The following keys are not loaded from the original tree after applying specified transforms:

{model layer printout removed}

I0920 10:00:43.076870 139574753909888 checkpointer.py:318] Finished restoring checkpoint in 3.04 seconds from gs://cantonese_llm/qwen3_maxtext_ckpt/Qwen3-30B-A3B-Base/0/items.
E0920 10:12:44.537173 17659 coordination_service.cc:1893] Use error polling to propagate the following error to all tasks: UNAVAILABLE: The following tasks are unhealthy (stopped sending heartbeats):
/job:jax_worker/replica:0/task:7
/job:jax_worker/replica:0/task:3
/job:jax_worker/replica:0/task:1
The tasks have crashed. Check the task logs for an earlier error, or scheduler events (e.g. preemption, eviction) to debug further. [type.googleapis.com/tensorflow.CoordinationServiceError='']
E0920 10:12:44.610459 17682 coordination_service_agent.cc:332] Polled an error from coordination service (this can be an error from this or another task).
F0920 10:12:44.654663 17682 client.h:75] Terminating process because the JAX distributed service detected fatal errors. This most likely indicates that another task died; see the other task logs for more details. Disable Python buffering, i.e. `python -u`, to be sure to see all the previous output. absl::Status: UNAVAILABLE: The following tasks are unhealthy (stopped sending heartbeats):
/job:jax_worker/replica:0/task:7
/job:jax_worker/replica:0/task:3
/job:jax_worker/replica:0/task:1
The tasks have crashed. Check the task logs for an earlier error, or scheduler events (e.g. preemption, eviction) to debug further.

RPC: /tensorflow.CoordinationService/PollForError [type.googleapis.com/tensorflow.CoordinationServiceError='']
/home/jed351/.local/share/uv/python/cpython-3.12.11-linux-x86_64-gnu/lib/python3.12/multiprocessing/resource_tracker.py:279: UserWarning: resource_tracker: There appear to be 16 leaked semaphore objects to clean up at shutdown
warnings.warn('resource_tracker: There appear to be %d '

Image

### Environment Information

TPU Creation
```bash
export TPU_PREFIX=node-2 # Use new names when you create new TPUs
export QR_ID=$TPU_PREFIX # Convenient to reuse the node names, but can be different

gcloud alpha compute tpus tpu-vm create $QR_ID \
--zone=europe-west4-a \
--accelerator-type=v6e-32 \
--version=v2-alpha-tpuv6e \
--spot \
--service-account=${tpu_service_account}

```

### Additional Context

_No response_

Guia de contribuição

Abrir o guia de contribuição

Avaliação

Esta issue ainda não foi avaliada.

Receba novas issues na sua caixa de entrada

Um resumo curto de issues do GitHub para quem está começando.