From 7047afafece5fdd77000afc1b15e2182ce321461 Mon Sep 17 00:00:00 2001 From: Rockdu <89372739+Rockdu@users.noreply.github.com> Date: Mon, 6 Jul 2026 17:56:52 -0700 Subject: [PATCH] [diffusion] fix: slice img_shapes per-sample in rollout response extractor (#29989) --- .../entrypoints/post_training/rollout_api.py | 20 +++++++++++++++---- 1 file changed, 16 insertions(+), 4 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/post_training/rollout_api.py b/python/sglang/multimodal_gen/runtime/entrypoints/post_training/rollout_api.py index 97e0dcb5b..159692d44 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/post_training/rollout_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/post_training/rollout_api.py @@ -33,21 +33,33 @@ logger = init_logger(__name__) router = APIRouter(prefix="/rollout", tags=["rollout"]) -def _extract_single_sample_tensor(obj: Any, sample_idx: int, batch_size: int) -> Any: +def _extract_single_sample_tensor( + obj: Any, sample_idx: int, batch_size: int, *, current_key: str | None = None +) -> Any: if isinstance(obj, torch.Tensor): if obj.dim() >= 1 and obj.shape[0] == batch_size: return obj[sample_idx].contiguous() return obj if isinstance(obj, dict): return { - k: _extract_single_sample_tensor(v, sample_idx, batch_size) + k: _extract_single_sample_tensor(v, sample_idx, batch_size, current_key=k) for k, v in obj.items() } if isinstance(obj, list): - return [_extract_single_sample_tensor(v, sample_idx, batch_size) for v in obj] + if current_key == "img_shapes" and len(obj) == batch_size: + return [obj[sample_idx]] + return [ + _extract_single_sample_tensor( + v, sample_idx, batch_size, current_key=current_key + ) + for v in obj + ] if isinstance(obj, tuple): return tuple( - _extract_single_sample_tensor(v, sample_idx, batch_size) for v in obj + _extract_single_sample_tensor( + v, sample_idx, batch_size, current_key=current_key + ) + for v in obj ) return obj