Skip to content

[RFC] Mixed-resolution rollouts, long-tail scheduling, sibling batching, and train/inference consistency checks #537

Description

@nussejzz

Status: Draft
Scope: RewardStack (unirl/reward/stack.py), the vLLM-Omni BAGEL adapter and worker pipeline (unirl/rollout/engine/vllm_omni), BAGEL replay (unirl/models/bagel/diffusion.py), TrainStack, Sample / Part, and later the trainside rollout engine

Summary

Follow-up to #483 (items 3 and 4) now that #445 has put the micro-batch loop on the worker. Four modules, each opened with one small PR; the RFC is deliberately incomplete and will grow as the PRs land.

  1. Mixed-resolution rollouts. Let every prompt group in one rollout have its own output canvas. Images first (BAGEL through vLLM-Omni); video and other engines later.
  2. Long-tail scheduling. Once groups cost different amounts, ranks finish at different times. A per-rank deque of rollout units ordered by a fitted cost, where an idle rank steals from the back of the busiest rank's deque, with results returned to the owner's row positions. Ring assignment is one policy on the same mechanism.
  3. Train/inference consistency checks. The rollout worker and the trainer each build the token sequence on their own and nothing compares the two. Start with a token-layout fingerprint that fails the replay on a mismatch, plus a first-update ratio guard.
  4. Sibling batching and prefix reuse inside a micro. Siblings of one group share prompt, source image and canvas, so they are equal-length: batch them densely and build the shared prefix once per group.

Without 1 the framework cannot run a long-tail experiment at all; 2 is the mitigation; 3 is what tells us a run is silently off-policy; 4 sets the unit cost that 2 schedules.

Motivation

  • One canvas per Part. sampling.height/width is a single value for the whole rollout, LatentSegment.latents is dense [N, K, seq, C], and the BAGEL adapter emits one StageSampling per call. A dataset that mixes 512 / 768 / 1024 square outputs (1024 / 2304 / 4096 image tokens for BAGEL) cannot be expressed, so every recipe resizes to one shape. fix(rollout): validate aligned prompt-group chunk boundaries #485 reaches the same conclusion from the trainside engine: heterogeneous-shape scheduling waits on per-row geometry.
  • The tail is shape-driven and known in advance. Per-step cost grows between linearly and quadratically in tokens, i.e. 4x to 16x from 512 to 1024. DP_SCATTER hands each rank a fixed set of prompt groups, so a rollout ends when the rank with the largest shapes ends. A list-scheduling simulation of 8 ranks x 4 groups with cost proportional to tokens gives a makespan of 10.0 units as dispatched today, 8.0 with whole-group stealing and 7.25 with half-group units, against an ideal of 6.28.
  • No layout comparison between worker and trainer. For BAGEL the worker and the trainer both call prepare_vae_latent and rebuild the prompt / source-image context independently. A different source-image resize, prompt normalization or upstream vLLM-Omni change shifts KV length or RoPE offsets; replay still runs, the ratio is just wrong. Ratio statistics are logged but nothing alerts on them. RFC: UniRL Experimental Plugin Framework for Train–Inference Parity #431 frames the same requirement for AR and lists diffusion as future scope.
  • Siblings are not batched everywhere. vLLM-Omni BAGEL packs a t2i group into one request and prefills the prompt once, but it2i sends one request per sample and rebuilds the three contexts (including the source image's VAE + ViT encode) for each. Trainside BAGEL recipes run forward_batch_size: 1; the packed path from perf(bagel): batch trainside UniGRPO rollout forwards #253 goes through the varlen attention interface, which without flash_attn is a Python loop of one SDPA call per sequence, and the prefix KV is physically copied per sibling.

Proposal

1. Mixed-resolution rollouts: per-group canvas

  • The canvas is row data, not config: metadata: {"canvas": [H, W]} on the dataset row, read through one helper Sample.canvases() that falls back to sampling.height/width. No new config key; a dataset without the field is byte-identical to today.
  • sampling.height/width stays the default canvas and becomes the storage bound: a row canvas may not exceed it in tokens. Trajectories of a smaller canvas are right-padded on the token axis to the default and trimmed by image_shape in replay, so Batch, the transport and the segment types do not change. A ragged latent segment is the long-term answer (see open questions) but touches every dense model.
  • First PR, vLLM-Omni BAGEL only: the input adapter splits the frontier into contiguous same-canvas runs and emits one GenerateCall per run; the worker already reads H/W per request and regenerates x_T per request from the driver's recipe, so it needs no change. Engines that cannot honor a row canvas reject it by name instead of ignoring it.
  • A group has one canvas, so GRPO advantages, RewardStack micro bounds and the reward scorers are unaffected.

2. Long-tail scheduling: cost-ordered deques with stealing

  • Unit. A chunk of one prompt group of at most micro_batch_size rows (a whole group for engines that pack groups). Every rank derives the same unit list from a group table, so a unit is addressed by id and nothing about it is negotiated at run time.
  • Deque. Each rank owns the units of its groups, sorted by estimated cost, largest first. The owner pops from the front; a rank whose deque is empty steals from the back (the smallest unit) of the rank with the most queued cost. Only units that have not started move. A coordinator actor holds the deques; workers pull from it between units, so no worker is ever pushed to and none waits on another.
  • Returning results to the owner. A helper returns only the frontier Part it produced. The driver concatenates what every rank returns and restores the request's row order by sample_id before advantages are computed, because Part.compute_advantages needs siblings contiguous. feat(trainer): opt-in row-level reward dispatch for diffusion #517 has the same restore-and-verify step for row-level reward dispatch; the two should share it.
  • Prefill. A helper recomputes the unit's context from the prompt and source image. Shipping KV across ranks costs more than the prefill it saves at these sizes and is out of scope.
  • Cost model. t = steps * (a*L + b*L^2 + c) + (d*L_src + e*L_src^2) + f per request, fitted per hardware and engine by least squares over a shape sweep (the way TeaCache ships fitted polynomial coefficients), stored in the recipe under reward_stack.cost, and corrected online by a per-rank measured/predicted ratio. Without coefficients the order falls back to token count.
  • Config. One key, reward_stack.schedule: owner | ring | steal; absent means today's DP_SCATTER path. ring rotates the chunks of each group across ranks so every rank runs one chunk of every group; it is exact when samples_per_prompt / micro_batch_size >= dp and needs no coordinator.
  • Reproducibility. Stealing changes which rank generates a row. Engines whose per-row noise is keyed by sample id (vLLM-Omni with the driver x_T recipe) are placement-free; trainside engines that draw from the global CUDA RNG are not, and steal is rejected there until they are.
  • Split: (2a) unit table, owner policy and row-order restore, bit-identical to DP_SCATTER at eta=0; (2b) coordinator, cost model and steal; (2c) ring.

3. Train/inference consistency: layout fingerprint and ratio guard

4. Sibling batching and prefix reuse inside a micro

Siblings are equal-length by construction, so a dense batch with standard attention is enough; varlen is only needed across groups.

Engine Sibling batch today Prefix reuse today
vLLM-Omni BAGEL t2i, cfg <= 1 one request per group prompt prefilled once, KV replicated per sibling
vLLM-Omni BAGEL it2i one request per sample none: contexts and source-image encode rebuilt per sample
trainside BAGEL recipes run bs=1; packed path is a per-sequence attention loop without flash_attn t2i context LRU keyed by prompt; it2i none; KV copied per sibling
trainside SD3 dense batch unique prompts encoded once, index_select
trainside Qwen-Image dense, recipes run bs=1 none

Related work

Module Related
1 per-group canvas #483 item 4, #485 (group-boundary validation; defers heterogeneous shapes to per-row geometry), #399 (ragged ImageSets), #470 (multi-image inputs), #253 (rejects mixed shapes in one LatentSegment)
2 scheduling #445 (RewardStack, merged), #517 (row-level reward dispatch), #362 (async rollout lanes), #289 (dynamic per-prompt dispatch), #367 (adaptive Ulysses SP), #484 (trace-driven rollout replay, a way to validate the cost model)
3 consistency #431 (parity plugin framework RFC), #430 / #504 / #505 (Qwen3-MoE parity), #533 / #536 (rollout vs trainer log-prob constants)
4 sibling batching #253 (packed trainside BAGEL forwards, merged), #529 (sibling group-affinity routing for AR), #443 / #449 (conditioning cache), #500 / #394 (batched replay timesteps)

Outside UniRL: KnapFormer (arXiv 2508.06001) balances mixed-resolution DiT training with a fitted FLOP cost model and reports a 17x load ratio before balancing; AdaptiveLoad (2605.17923) finds step latency correlates 0.35 with token count and 0.92 with squared length. On the LLM side RollPacker (2509.21009), Seer (2511.14617) and Laminar (2510.12633) move work off slow replicas, APRIL (2509.18521) over-provisions and truncates, TailSieve (2608.22788) isolates tail prompts. Diffusion cost is known before generation, so ordering and stealing are enough here and length prediction or partial rollouts are not needed.

Goals / non-goals

Suggested order

  1. 3: layout fingerprint + ratio guard (independent, and lets the following PRs quote "fingerprints match" in their test plans).
  2. 1: per-group canvas for vLLM-Omni BAGEL.
  3. 4a: it2i sibling packing, on top of 1 (same adapter function).
  4. 2a: unit table and row-order restore.
  5. 2b: cost model and stealing, measured on a mixed-resolution dataset from 1; 2c: ring.
  6. 4b / 4c, video, other engines.

Open questions

  • How should the diffusion checks plug into the RFC: UniRL Experimental Plugin Framework for Train–Inference Parity #431 parity framework and rl-kernel: as a provider there, or as a model-side contract the framework calls?
  • Pad-and-trim versus a ragged latent segment: at what spread of canvases does the padding in storage and transport justify the larger refactor, and should it share a primitive with [RFC] Ordered Ragged ImageSet / ImageSets Primitives #399?
  • Under a schedule, should overlap default to on, and should owner eventually replace the DP_SCATTER path so there is one code path?
  • Can trainside engines be made placement-free (per-row generators instead of the global RNG) so that steal applies to them?
  • Should the fitted cost coefficients live in the recipe, or in a per-hardware file the recipe points at?

AI assistance: this RFC was drafted with Claude Code from our measurements, a read of the current code and the linked PRs; I reviewed it.

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

    enhancementNew feature or request

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions