Bug report
Summary
python3 -m maxtext.trainers.post_train.rl.train_rl fails during vLLM engine
initialization (_precompile_gather_logprobs) with:
jax._src.core.ShardingTypeError: The input should be unsharded over the axis along which to
compute the top_k values. Got input type=float32[8@data,128256@model] and axis=1
The call goes through the _compat_top_k shim added in #4708
(src/maxtext/integration/tunix/tunix_adapter.py). It looks like the shim's
fallback never reshards here, because the operand is a tracer inside jax.jit.
Reproduction
uv venv --python 3.12 --seed maxtext_venv && source maxtext_venv/bin/activate
UV_TORCH_BACKEND=cpu uv pip install "maxtext[tpu-post-train]==0.2.4" --resolution=lowest
install_tpu_post_train_extra_deps
python3 -m maxtext.checkpoint_conversion.to_maxtext \
model_name=llama3.1-8b-Instruct hf_access_token=$HF_TOKEN \
base_output_directory=/dev/shm/llama3.1-8b-Instruct/mt-format/ \
scan_layers=True use_multimodal=False hardware=cpu skip_jax_distributed_system=true \
checkpoint_storage_use_zarr3=1 checkpoint_storage_use_ocdbt=1 --lazy_load_tensors=False
python3 -m maxtext.trainers.post_train.rl.train_rl \
model_name=llama3.1-8b-Instruct \
load_parameters_path=/dev/shm/llama3.1-8b-Instruct/mt-format/0/items \
run_name=test base_output_directory=/dev/shm/llama3.1-8b-Instruct/post-train/ \
chips_per_vm=8 num_batches=50 num_test_batches=10 \
rollout_data_parallelism=1 rollout_tensor_parallelism=8
(rollout_tensor_parallelism=-1 hits the same code path; the run above used 8.)
Relevant resolved config: shard_mode: AUTO, ici_fsdp_parallelism: -1,
max_logprobs: 1, logprobs_mode: processed_logprobs.
Traceback (trimmed)
File ".../tunix/rl/rl_cluster.py", line 487, in _init_engine
self._rollout = self.cluster_config.rollout_engine(
File ".../maxtext/integration/vllm/maxtext_vllm_rollout.py", line 587, in __init__
self._sampler = MaxTextVllmSampler(
File ".../tunix/generate/vllm_sampler.py", line 161, in __init__
self.llm = LLM(**self.args)
...
File ".../vllm/v1/engine/core.py", line 324, in _initialize_kv_caches
self.model_executor.initialize_from_config(kv_cache_configs)
...
File ".../tpu_inference/worker/tpu_worker.py", line 710, in initialize_from_config
self.model_runner.compilation_manager._precompile_gather_logprobs()
File ".../tpu_inference/runner/compilation_manager.py", line 166, in _run_compilation
lowered = fn.lower(*args, **call_kwargs)
File ".../tpu_inference/layers/jax/sample/sampling.py", line 162, in compute_and_gather_logprobs
return gather_logprobs(logprobs, next_tokens, max_logprobs)
File ".../tpu_inference/layers/jax/sample/sampling.py", line 262, in gather_logprobs
topk_logprobs, topk_indices = jax.lax.top_k(logprobs, k=num_logprobs)
File ".../maxtext/integration/tunix/tunix_adapter.py", line 64, in _compat_top_k
return _orig_top_k(operand, k, axis=axis)
jax._src.core.ShardingTypeError: The input should be unsharded over the axis along which to
compute the top_k values. Got input type=float32[8@data,128256@model] and axis=1
Analysis
compute_and_gather_logprobs is decorated with @jax.jit, so the shim receives a
tracer, not a concrete jax.Array. The fallback reads the sharding with:
sharding = getattr(operand, "sharding", None)
if sharding is not None and hasattr(sharding, "spec") and hasattr(sharding, "mesh"):
... # reshard
return _orig_top_k(operand, k, axis=axis)
I think getattr(..., "sharding", None) returns None for the tracer, so the
reshard branch is skipped. The final line (line 64 in the traceback) then
re-raises the same error. The tracer's sharding is available on its abstract
value (jax.typeof(operand).sharding), and the error message itself prints it
(8@data,128256@model).
Logs/Output
No response
Environment Information
| Component |
Version |
| MaxText |
maxtext[tpu-post-train]==0.2.4 (PyPI), plus install_tpu_post_train_extra_deps |
| jax / jaxlib |
0.11.0 |
| libtpu |
0.0.44 |
| flax |
0.12.8 |
| google-tunix |
c4ec573d29e4c3a3955b348256d464b119c8a6d1 (pinned by 0.2.4) |
| tpu-inference |
7ecc401e6faefe3f793f3e4a8e6e2dbc49eca868 (pinned by 0.2.4) |
| vllm |
0ba2aa35a81dcc3246b26291368b53fa2389c7d7 (pinned by 0.2.4) |
| Python |
3.12.14 (uv) |
| Hardware |
TPU v6e-8, single host (ct6e-standard-8t) |
| OS image |
ubuntu-accel-2204-amd64-tpu-v5e-v5p-v6e |
I also checked current main (41389c2), and the code involved looks the same:
_compat_top_k is unchanged. tpu-inference b67ae5f (pinned by main) and
tpu-inference main still call jax.lax.top_k(logprobs, ...) in
gather_logprobs without resharding first.
Additional Context
No response
Bug report
Summary
python3 -m maxtext.trainers.post_train.rl.train_rlfails during vLLM engineinitialization (
_precompile_gather_logprobs) with:The call goes through the
_compat_top_kshim added in #4708(
src/maxtext/integration/tunix/tunix_adapter.py). It looks like the shim'sfallback never reshards here, because the operand is a tracer inside
jax.jit.Reproduction
(
rollout_tensor_parallelism=-1hits the same code path; the run above used8.)Relevant resolved config:
shard_mode: AUTO,ici_fsdp_parallelism: -1,max_logprobs: 1,logprobs_mode: processed_logprobs.Traceback (trimmed)
Analysis
compute_and_gather_logprobsis decorated with@jax.jit, so the shim receives atracer, not a concrete
jax.Array. The fallback reads the sharding with:I think
getattr(..., "sharding", None)returnsNonefor the tracer, so thereshard branch is skipped. The final line (line 64 in the traceback) then
re-raises the same error. The tracer's sharding is available on its abstract
value (
jax.typeof(operand).sharding), and the error message itself prints it(
8@data,128256@model).Logs/Output
No response
Environment Information
maxtext[tpu-post-train]==0.2.4(PyPI), plusinstall_tpu_post_train_extra_depsc4ec573d29e4c3a3955b348256d464b119c8a6d1(pinned by 0.2.4)7ecc401e6faefe3f793f3e4a8e6e2dbc49eca868(pinned by 0.2.4)0ba2aa35a81dcc3246b26291368b53fa2389c7d7(pinned by 0.2.4)ct6e-standard-8t)ubuntu-accel-2204-amd64-tpu-v5e-v5p-v6eI also checked current
main(41389c2), and the code involved looks the same:_compat_top_kis unchanged. tpu-inferenceb67ae5f(pinned bymain) andtpu-inference
mainstill calljax.lax.top_k(logprobs, ...)ingather_logprobswithout resharding first.Additional Context
No response