2828import logging
2929import time
3030import jax
31+ import jax .numpy as jnp
3132from flax import nnx
33+ from flax .traverse_util import flatten_dict , unflatten_dict
34+
3235from pathwaysutils .experimental import reshard as _experimental_reshard
3336from tunix .generate import mappings
3437from tunix .generate .vllm_sampler import VllmConfig , VllmSampler
@@ -45,7 +48,83 @@ def _create_model_converter(model_name: str, config: Any, mesh: jax.sharding.Mes
4548 elif model_name in {"qwen3.5-35b-a3b" }:
4649 return Qwen35MaxTextToVLLMConverter (config = config , mesh = mesh )
4750
48- raise ValueError (f"No MaxText->vLLM converter registered for model { model_name !r} ." )
51+ return None
52+
53+
54+ def unroll_gemma_scanned_weights (weights ):
55+ """Workaround for tunix unstacking bug with Gemma 3/4 scanned blocks.
56+
57+ tunix fails to map nested layers like `layers.layers_0` to `layers_X`.
58+ We manually unroll them here if we detect the structure.
59+ """
60+ if not hasattr (weights , "to_pure_dict" ):
61+ return weights
62+
63+ flat_w = flatten_dict (weights .to_pure_dict (), sep = "/" )
64+ new_flat_w = {}
65+
66+ # Check if this is actually a scanned Gemma 3/4 checkpoint
67+ # by looking for the scanned nested structure.
68+ is_gemma_scanned = any ("decoder/layers/layers_0/" in k or "decoder/scanned_blocks/layers_0/" in k for k in flat_w )
69+
70+ if not is_gemma_scanned :
71+ return weights
72+
73+ logging .info ("MaxTextVllmSampler: Detected Gemma scanned weights structure. Unrolling along axis 1..." )
74+
75+ # Determine attention pattern length and scan length
76+ pattern_keys = set ()
77+ scan_length = 0
78+ for k , v in flat_w .items ():
79+ if "decoder/layers/layers_" in k or "decoder/scanned_blocks/layers_" in k :
80+ layer_sub_idx = k .split ("layers_" )[- 1 ].split ("/" )[0 ]
81+ pattern_keys .add (int (layer_sub_idx ))
82+ # In MaxText, Gemma uses param_scan_axis=1, so the scan dimension is at axis 1
83+ if hasattr (v , "shape" ) and len (v .shape ) > 1 :
84+ scan_length = max (scan_length , v .shape [1 ])
85+
86+ pattern_length = max (pattern_keys ) + 1 if pattern_keys else 0
87+ logging .info ("MaxTextVllmSampler: Discovered scan_length=%d, pattern_length=%d" , scan_length , pattern_length )
88+
89+ unrolled_count = 0
90+ for k , v in flat_w .items ():
91+ if "decoder/layers/layers_" in k or "decoder/scanned_blocks/layers_" in k :
92+ # Unstack the array along the 1st axis
93+ if "decoder/scanned_blocks/layers_" in k :
94+ parts = k .split ("decoder/scanned_blocks/layers_" )
95+ else :
96+ parts = k .split ("decoder/layers/layers_" )
97+
98+ layer_sub_idx = int (parts [1 ].split ("/" )[0 ])
99+ suffix = "/" + "/" .join (parts [1 ].split ("/" )[1 :])
100+
101+ if hasattr (v , "shape" ) and len (v .shape ) > 1 :
102+ v_swapped = jnp .swapaxes (v , 1 , 0 )
103+ unstacked = [v_swapped [i ] for i in range (scan_length )]
104+ else :
105+ unstacked = [v ] * scan_length
106+
107+ for i in range (scan_length ):
108+ global_idx = i * pattern_length + layer_sub_idx
109+ # Map back to nnx.List format which uses layers/X/ instead of layers_X
110+ new_flat_w [f"decoder/layers/{ global_idx } { suffix } " ] = unstacked [i ]
111+ unrolled_count += 1
112+
113+ elif "decoder/layers_remainder/layers_" in k :
114+ layer_sub_idx = int (k .split ("decoder/layers_remainder/layers_" )[1 ].split ("/" )[0 ])
115+ suffix = "/" + "/" .join (k .split ("decoder/layers_remainder/layers_" )[1 ].split ("/" )[1 :])
116+
117+ global_idx = scan_length * pattern_length + layer_sub_idx
118+ new_flat_w [f"decoder/layers/{ global_idx } { suffix } " ] = v
119+ unrolled_count += 1
120+ else :
121+ new_flat_w [k ] = v
122+
123+ logging .info (
124+ "MaxTextVllmSampler: Successfully unrolled %d scanned tensor components into vLLM-compatible nnx.List format." ,
125+ unrolled_count ,
126+ )
127+ return unflatten_dict (new_flat_w , sep = "/" )
49128
50129
51130class MaxTextVllmSampler (VllmSampler ):
@@ -73,6 +152,12 @@ def update_params(
73152 ):
74153 """Update the vLLM runner weights from a MaxText state tree."""
75154 if self ._converter is None :
155+ # --- Workaround for tunix unstacking bug with Gemma 3/4 scanned blocks ---
156+ # tunix fails to map nested layers like `layers.layers_0` to `layers_X`.
157+ # We manually unroll them here if we detect the structure.
158+ updated_weights = unroll_gemma_scanned_weights (updated_weights )
159+ # --- End Workaround ---
160+
76161 super ().update_params (updated_weights , filter_types )
77162 return None
78163
@@ -182,6 +267,19 @@ def __init__(
182267 model = rollout_actor ,
183268 backend = "vllm_jax" ,
184269 )
270+ engine_kwargs = {
271+ "max_model_len" : cache_config_or_size ,
272+ "model" : rollout_config .rollout_vllm_model_version ,
273+ "swap_space" : getattr (rollout_config , "rollout_vllm_swap_space_size_gb" , maxtext_config .swap_space_vllm_gb ),
274+ # Async scheduling causes KeyError in dp_scheduler on slow models
275+ # (30B+) where inference latency exceeds the scheduler's window.
276+ "async_scheduling" : rollout_config .rollout_vllm_async_scheduling ,
277+ }
278+
279+ # Merge additional kwargs like dtype and hf_overrides provided by train_rl.py
280+ if hasattr (rollout_config , "rollout_vllm_kwargs" ) and rollout_config .rollout_vllm_kwargs :
281+ engine_kwargs .update (rollout_config .rollout_vllm_kwargs )
282+
185283 self ._sampler = MaxTextVllmSampler (
186284 tokenizer = tokenizer ,
187285 config = VllmConfig ( # pylint: disable=unexpected-keyword-arg,no-value-for-parameter
@@ -195,14 +293,8 @@ def __init__(
195293 tensor_parallel_size = rollout_config .tensor_parallel_size ,
196294 data_parallel_size = rollout_config .data_parallel_size ,
197295 enable_dp_attention = rollout_config .rollout_vllm_enable_dp_attention ,
198- engine_kwargs = {
199- "max_model_len" : cache_config_or_size ,
200- "model" : rollout_config .rollout_vllm_model_version ,
201- "swap_space" : rollout_config .rollout_vllm_swap_space_size_gb ,
202- # Async scheduling causes KeyError in dp_scheduler on slow models
203- # (30B+) where inference latency exceeds the scheduler's window.
204- "async_scheduling" : rollout_config .rollout_vllm_async_scheduling ,
205- },
296+ engine_kwargs = engine_kwargs ,
297+ additional_config = getattr (rollout_config , "rollout_vllm_additional_config" , None ),
206298 ),
207299 converter = converter ,
208300 )
0 commit comments