Skip to content

[Bug] _compat_top_k shim does not handle traced inputs → ShardingTypeError in tpu-inference gather_logprobs on JAX 0.11 #5440

Description

@piowag

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

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    bugSomething isn't working

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions