From b66085a305002574fe22301f00419b56a93da6f5 Mon Sep 17 00:00:00 2001 From: zhouzhuoxin Date: Thu, 20 Aug 2026 11:13:48 +0800 Subject: [PATCH 1/3] perf(bagel): pack FlowGRPO replay steps Pack no-CFG T2I/IT2I replay steps as isolated varlen sequences to reduce repeated FSDP forwards. --- .../diffusion/bagel/bagel_editreward.yaml | 1 + .../diffusion/bagel/bagel_it2i_vllmomni.yaml | 1 + .../diffusion/bagel/bagel_trainside_lora.yaml | 1 + examples/diffusion/bagel/bagel_vllmomni.yaml | 1 + .../diffusion/bagel/bagel_vllmomni_async.yaml | 1 + unirl/models/bagel/config.py | 1 + unirl/models/bagel/diffusion.py | 69 ++++++++++++++++++- unirl/models/bagel/pipeline.py | 1 + unirl/models/bagel/rl_ops.py | 18 +++++ .../bagel/vendor/modeling/bagel/bagel.py | 5 +- 10 files changed, 95 insertions(+), 4 deletions(-) diff --git a/examples/diffusion/bagel/bagel_editreward.yaml b/examples/diffusion/bagel/bagel_editreward.yaml index 4dc77695a..7028684ba 100644 --- a/examples/diffusion/bagel/bagel_editreward.yaml +++ b/examples/diffusion/bagel/bagel_editreward.yaml @@ -72,6 +72,7 @@ bundle: shift: 3.0 use_lora: true enable_vit: true # EDIT: load the SigLIP ViT/und path for source-image conditioning + replay_step_pack_size: 3 pipeline: _target_: unirl.models.bagel.pipeline.BagelPipeline diff --git a/examples/diffusion/bagel/bagel_it2i_vllmomni.yaml b/examples/diffusion/bagel/bagel_it2i_vllmomni.yaml index 52a87ab6b..71a243fe4 100644 --- a/examples/diffusion/bagel/bagel_it2i_vllmomni.yaml +++ b/examples/diffusion/bagel/bagel_it2i_vllmomni.yaml @@ -94,6 +94,7 @@ bundle: shift: 3.0 use_lora: true # read by the engine's WeightSync (uses_lora) enable_vit: true # EDIT: the und ViT the replay context rebuild needs + replay_step_pack_size: 3 pipeline: _target_: unirl.models.bagel.pipeline.BagelPipeline diff --git a/examples/diffusion/bagel/bagel_trainside_lora.yaml b/examples/diffusion/bagel/bagel_trainside_lora.yaml index c8a1cd4e7..cd18c9368 100644 --- a/examples/diffusion/bagel/bagel_trainside_lora.yaml +++ b/examples/diffusion/bagel/bagel_trainside_lora.yaml @@ -43,6 +43,7 @@ bundle: model_precision: bf16 shift: 3.0 use_lora: true + replay_step_pack_size: 2 pipeline: _target_: unirl.models.bagel.pipeline.BagelPipeline diff --git a/examples/diffusion/bagel/bagel_vllmomni.yaml b/examples/diffusion/bagel/bagel_vllmomni.yaml index 11bfc3d0e..877aaf51e 100644 --- a/examples/diffusion/bagel/bagel_vllmomni.yaml +++ b/examples/diffusion/bagel/bagel_vllmomni.yaml @@ -55,6 +55,7 @@ bundle: model_precision: bf16 shift: 3.0 use_lora: true # read by the engine's WeightSync (uses_lora) + replay_step_pack_size: 2 pipeline: _target_: unirl.models.bagel.pipeline.BagelPipeline diff --git a/examples/diffusion/bagel/bagel_vllmomni_async.yaml b/examples/diffusion/bagel/bagel_vllmomni_async.yaml index a2227d0ef..5fd978673 100644 --- a/examples/diffusion/bagel/bagel_vllmomni_async.yaml +++ b/examples/diffusion/bagel/bagel_vllmomni_async.yaml @@ -43,6 +43,7 @@ bundle: model_precision: bf16 shift: 3.0 use_lora: true # read by the engine's WeightSync (uses_lora) + replay_step_pack_size: 2 pipeline: _target_: unirl.models.bagel.pipeline.BagelPipeline diff --git a/unirl/models/bagel/config.py b/unirl/models/bagel/config.py index 57c547d2a..93a06f6d6 100644 --- a/unirl/models/bagel/config.py +++ b/unirl/models/bagel/config.py @@ -57,6 +57,7 @@ class BagelPipelineConfig: cache_t2i_contexts: bool = True context_cache_size: int = 32 + replay_step_pack_size: int = 1 def __post_init__(self) -> None: validate_precision_type(self.model_precision, field="BagelPipelineConfig.model_precision") diff --git a/unirl/models/bagel/diffusion.py b/unirl/models/bagel/diffusion.py index 523c1f455..7792774f2 100644 --- a/unirl/models/bagel/diffusion.py +++ b/unirl/models/bagel/diffusion.py @@ -160,6 +160,7 @@ def __init__( autocast_precision: str = "bf16", trajectory_precision: str = "fp32", logprob_precision: str = "fp32", + replay_step_pack_size: int = 1, ) -> None: self.model = model self.step = step if step is not None else BagelDiffusionStep() @@ -167,6 +168,8 @@ def __init__( self.autocast_dtype = parse_torch_dtype(autocast_precision, field_name="autocast_precision") self.trajectory_dtype = parse_torch_dtype(trajectory_precision, field_name="trajectory_precision") self.logprob_dtype = parse_torch_dtype(logprob_precision, field_name="logprob_precision") + self.replay_step_pack_size = int(replay_step_pack_size) + require(self.replay_step_pack_size >= 1, "replay_step_pack_size must be >= 1") def _autocast_ctx(self, device: torch.device): if device.type == "cuda" and self.autocast_dtype in (torch.float16, torch.bfloat16): @@ -265,14 +268,18 @@ def _build_generation_inputs( gi = bagel.prepare_vae_latent( curr_kvlens=gen["kv_lens"], curr_rope=gen["ropes"], - image_sizes=[image_shape], + image_sizes=[image_shape] * len(gen["kv_lens"]), new_token_ids=self.model.new_token_ids, ) gi_cfg_text = bagel.prepare_vae_latent_cfg( - curr_kvlens=cfg_text["kv_lens"], curr_rope=cfg_text["ropes"], image_sizes=[image_shape] + curr_kvlens=cfg_text["kv_lens"], + curr_rope=cfg_text["ropes"], + image_sizes=[image_shape] * len(cfg_text["kv_lens"]), ) gi_cfg_img = bagel.prepare_vae_latent_cfg( - curr_kvlens=cfg_img["kv_lens"], curr_rope=cfg_img["ropes"], image_sizes=[image_shape] + curr_kvlens=cfg_img["kv_lens"], + curr_rope=cfg_img["ropes"], + image_sizes=[image_shape] * len(cfg_img["kv_lens"]), ) return _to_device(gi, device), _to_device(gi_cfg_text, device), _to_device(gi_cfg_img, device) @@ -436,6 +443,62 @@ def replay( conditions, differentiable=torch.is_grad_enabled(), ) + if self.replay_step_pack_size > 1: + require( + float(params.cfg_text_scale) == 1.0 and float(params.cfg_img_scale) == 1.0, + "Bagel packed replay only supports the FlowGRPO no-CFG path", + ) + log_prob_chunks: List[torch.Tensor] = [] + prev_mean_chunks: List[torch.Tensor] = [] + with self._autocast_ctx(device): + for start in range(0, len(target), self.replay_step_pack_size): + chunk = target[start : start + self.replay_step_pack_size] + repeats = len(chunk) + packed_gen = rl_ops.repeat_context(gen, repeats) + gi, gi_cfg_text, gi_cfg_img = self._build_generation_inputs( + packed_gen, cfg_text, cfg_img, image_shape, device=device + ) + forward_kwargs = self._forward_kwargs( + packed_gen, cfg_text, cfg_img, gi, gi_cfg_text, gi_cfg_img, params + ) + + x_t = torch.stack([segment.latents_at(i)[0].to(device) for i in chunk]) + prev_sample = torch.stack([segment.latents_at(i + 1)[0].to(device) for i in chunk]) + t_cur = torch.stack([schedule[i] for i in chunk]) + t_next = torch.stack([schedule[i + 1] for i in chunk]) + seq = int(x_t.shape[1]) + timestep = t_cur[:, None].expand(repeats, seq).reshape(-1) + + rl_ops.disable_inference_cache(bagel) + v_t = rl_ops.forward_flow( + bagel, + x_t=x_t.flatten(0, 1), + timestep=timestep, + cfg_text_scale=1.0, + cfg_img_scale=1.0, + **forward_kwargs, + ).view_as(x_t) + _, log_prob, prev_mean = self.strategy.denoise( + noise_pred=v_t, + sample=x_t, + sigma=t_cur, + sigma_next=t_next, + eta=float(params.eta), + prev_sample=prev_sample, + sigma_max=float(sigma_max), + ) + if log_prob is None or prev_mean is None: + raise RuntimeError("Bagel packed replay requires a stochastic FlowGRPO SDE strategy") + log_prob_chunks.append(log_prob.reshape(-1)) + prev_mean_chunks.append(prev_mean) + + return ReplayResult( + log_probs=torch.cat(log_prob_chunks).unsqueeze(0).to(dtype=self.logprob_dtype), + prev_sample_means=torch.cat(prev_mean_chunks) + .unsqueeze(0) + .to(dtype=self.trajectory_dtype), + ) + gi, gi_cfg_text, gi_cfg_img = self._build_generation_inputs(gen, cfg_text, cfg_img, image_shape, device=device) forward_kwargs = self._forward_kwargs(gen, cfg_text, cfg_img, gi, gi_cfg_text, gi_cfg_img, params) diff --git a/unirl/models/bagel/pipeline.py b/unirl/models/bagel/pipeline.py index 94f16d211..3872a7aa5 100644 --- a/unirl/models/bagel/pipeline.py +++ b/unirl/models/bagel/pipeline.py @@ -74,6 +74,7 @@ def __init__( autocast_precision=autocast_precision, trajectory_precision=trajectory_precision, logprob_precision=logprob_precision, + replay_step_pack_size=_cfg_get(getattr(bundle, "config", None), "replay_step_pack_size", 1), ) self.diffusion = diffusion self.vae_decode = vae_decode if vae_decode is not None else BagelVAEDecodeStage(bundle) diff --git a/unirl/models/bagel/rl_ops.py b/unirl/models/bagel/rl_ops.py index 2571fc467..5912e0989 100644 --- a/unirl/models/bagel/rl_ops.py +++ b/unirl/models/bagel/rl_ops.py @@ -24,6 +24,7 @@ "prefill_text_split", "prefill_vit_split", "require_inference_dispatch", + "repeat_context", "resize_input_image", "score_response", "score_response_with_prompt", @@ -162,6 +163,23 @@ def clone_context(ctx: Dict[str, Any]) -> Dict[str, Any]: } +def repeat_context(ctx: Dict[str, Any], repeats: int) -> Dict[str, Any]: + """Repeat a packed KV context along its varlen sequence axis.""" + if repeats == 1: + return ctx + cache = ctx["past_key_values"] + repeated_cache = type(cache)(cache.num_layers) + for layer_idx in range(cache.num_layers): + key, value = cache.key_cache[layer_idx], cache.value_cache[layer_idx] + repeated_cache.key_cache[layer_idx] = None if key is None else torch.cat([key] * repeats, dim=0) + repeated_cache.value_cache[layer_idx] = None if value is None else torch.cat([value] * repeats, dim=0) + return { + "kv_lens": list(ctx["kv_lens"]) * repeats, + "ropes": list(ctx["ropes"]) * repeats, + "past_key_values": repeated_cache, + } + + def update_context_text( bundle: Any, text: str, diff --git a/unirl/models/bagel/vendor/modeling/bagel/bagel.py b/unirl/models/bagel/vendor/modeling/bagel/bagel.py index 314769a25..9d88a8816 100644 --- a/unirl/models/bagel/vendor/modeling/bagel/bagel.py +++ b/unirl/models/bagel/vendor/modeling/bagel/bagel.py @@ -797,7 +797,10 @@ def _forward_flow( packed_sequence = packed_text_embedding.new_zeros((sum(packed_seqlens), self.hidden_size)) packed_sequence[packed_text_indexes] = packed_text_embedding - assert timestep.unique().shape[0] == 1 + if timestep.ndim != 1 or timestep.shape[0] != x_t.shape[0]: + raise ValueError( + f"timestep must contain one value per latent token, got {tuple(timestep.shape)} for {x_t.shape[0]}" + ) packed_pos_embed = self.latent_pos_embed(packed_vae_position_ids) packed_timestep_embeds = self.time_embedder(timestep) x_t = self.vae2llm(x_t) + packed_timestep_embeds + packed_pos_embed From b0526ec3e7a3e7f82d97a0bd51b22b369f937b4d Mon Sep 17 00:00:00 2001 From: zhouzhuoxin Date: Thu, 20 Aug 2026 12:35:44 +0800 Subject: [PATCH 2/3] refactor(bagel): infer replay pack size Use pipeline.batch_replay_steps as the only opt-in and pack all runtime-resolved SDE targets. --- .../diffusion/bagel/bagel_editreward.yaml | 2 +- .../diffusion/bagel/bagel_it2i_vllmomni.yaml | 2 +- .../diffusion/bagel/bagel_trainside_lora.yaml | 2 +- examples/diffusion/bagel/bagel_vllmomni.yaml | 2 +- .../diffusion/bagel/bagel_vllmomni_async.yaml | 2 +- unirl/models/bagel/config.py | 1 - unirl/models/bagel/diffusion.py | 90 +++++++++---------- unirl/models/bagel/pipeline.py | 3 +- 8 files changed, 47 insertions(+), 57 deletions(-) diff --git a/examples/diffusion/bagel/bagel_editreward.yaml b/examples/diffusion/bagel/bagel_editreward.yaml index 7028684ba..eb0274be5 100644 --- a/examples/diffusion/bagel/bagel_editreward.yaml +++ b/examples/diffusion/bagel/bagel_editreward.yaml @@ -72,10 +72,10 @@ bundle: shift: 3.0 use_lora: true enable_vit: true # EDIT: load the SigLIP ViT/und path for source-image conditioning - replay_step_pack_size: 3 pipeline: _target_: unirl.models.bagel.pipeline.BagelPipeline + batch_replay_steps: true autocast_precision: bf16 trajectory_precision: fp32 logprob_precision: fp32 diff --git a/examples/diffusion/bagel/bagel_it2i_vllmomni.yaml b/examples/diffusion/bagel/bagel_it2i_vllmomni.yaml index 71a243fe4..005362609 100644 --- a/examples/diffusion/bagel/bagel_it2i_vllmomni.yaml +++ b/examples/diffusion/bagel/bagel_it2i_vllmomni.yaml @@ -94,10 +94,10 @@ bundle: shift: 3.0 use_lora: true # read by the engine's WeightSync (uses_lora) enable_vit: true # EDIT: the und ViT the replay context rebuild needs - replay_step_pack_size: 3 pipeline: _target_: unirl.models.bagel.pipeline.BagelPipeline + batch_replay_steps: true autocast_precision: bf16 trajectory_precision: fp32 logprob_precision: fp32 diff --git a/examples/diffusion/bagel/bagel_trainside_lora.yaml b/examples/diffusion/bagel/bagel_trainside_lora.yaml index cd18c9368..dc46ab402 100644 --- a/examples/diffusion/bagel/bagel_trainside_lora.yaml +++ b/examples/diffusion/bagel/bagel_trainside_lora.yaml @@ -43,10 +43,10 @@ bundle: model_precision: bf16 shift: 3.0 use_lora: true - replay_step_pack_size: 2 pipeline: _target_: unirl.models.bagel.pipeline.BagelPipeline + batch_replay_steps: true # v3-equivalent: bf16 heads + bf16 transformer compute under autocast bf16 (no # fp32-heads → vendor stays byte-pristine). fp32 LoRA master (below) is the only # precision lever — it is THE reward-collapse fix. diff --git a/examples/diffusion/bagel/bagel_vllmomni.yaml b/examples/diffusion/bagel/bagel_vllmomni.yaml index 877aaf51e..5db647259 100644 --- a/examples/diffusion/bagel/bagel_vllmomni.yaml +++ b/examples/diffusion/bagel/bagel_vllmomni.yaml @@ -55,10 +55,10 @@ bundle: model_precision: bf16 shift: 3.0 use_lora: true # read by the engine's WeightSync (uses_lora) - replay_step_pack_size: 2 pipeline: _target_: unirl.models.bagel.pipeline.BagelPipeline + batch_replay_steps: true autocast_precision: bf16 trajectory_precision: fp32 logprob_precision: fp32 diff --git a/examples/diffusion/bagel/bagel_vllmomni_async.yaml b/examples/diffusion/bagel/bagel_vllmomni_async.yaml index 5fd978673..9462e39ce 100644 --- a/examples/diffusion/bagel/bagel_vllmomni_async.yaml +++ b/examples/diffusion/bagel/bagel_vllmomni_async.yaml @@ -43,10 +43,10 @@ bundle: model_precision: bf16 shift: 3.0 use_lora: true # read by the engine's WeightSync (uses_lora) - replay_step_pack_size: 2 pipeline: _target_: unirl.models.bagel.pipeline.BagelPipeline + batch_replay_steps: true autocast_precision: bf16 trajectory_precision: fp32 logprob_precision: fp32 diff --git a/unirl/models/bagel/config.py b/unirl/models/bagel/config.py index 93a06f6d6..57c547d2a 100644 --- a/unirl/models/bagel/config.py +++ b/unirl/models/bagel/config.py @@ -57,7 +57,6 @@ class BagelPipelineConfig: cache_t2i_contexts: bool = True context_cache_size: int = 32 - replay_step_pack_size: int = 1 def __post_init__(self) -> None: validate_precision_type(self.model_precision, field="BagelPipelineConfig.model_precision") diff --git a/unirl/models/bagel/diffusion.py b/unirl/models/bagel/diffusion.py index 7792774f2..f9c6bd83d 100644 --- a/unirl/models/bagel/diffusion.py +++ b/unirl/models/bagel/diffusion.py @@ -160,7 +160,7 @@ def __init__( autocast_precision: str = "bf16", trajectory_precision: str = "fp32", logprob_precision: str = "fp32", - replay_step_pack_size: int = 1, + batch_replay_steps: bool = False, ) -> None: self.model = model self.step = step if step is not None else BagelDiffusionStep() @@ -168,8 +168,8 @@ def __init__( self.autocast_dtype = parse_torch_dtype(autocast_precision, field_name="autocast_precision") self.trajectory_dtype = parse_torch_dtype(trajectory_precision, field_name="trajectory_precision") self.logprob_dtype = parse_torch_dtype(logprob_precision, field_name="logprob_precision") - self.replay_step_pack_size = int(replay_step_pack_size) - require(self.replay_step_pack_size >= 1, "replay_step_pack_size must be >= 1") + # This Bagel-specific opt-in intentionally retains rollout-sourced anchors. + self._batch_replay_steps = bool(batch_replay_steps) def _autocast_ctx(self, device: torch.device): if device.type == "cuda" and self.autocast_dtype in (torch.float16, torch.bfloat16): @@ -443,60 +443,50 @@ def replay( conditions, differentiable=torch.is_grad_enabled(), ) - if self.replay_step_pack_size > 1: + if self._batch_replay_steps and len(target) > 1: require( float(params.cfg_text_scale) == 1.0 and float(params.cfg_img_scale) == 1.0, "Bagel packed replay only supports the FlowGRPO no-CFG path", ) - log_prob_chunks: List[torch.Tensor] = [] - prev_mean_chunks: List[torch.Tensor] = [] - with self._autocast_ctx(device): - for start in range(0, len(target), self.replay_step_pack_size): - chunk = target[start : start + self.replay_step_pack_size] - repeats = len(chunk) - packed_gen = rl_ops.repeat_context(gen, repeats) - gi, gi_cfg_text, gi_cfg_img = self._build_generation_inputs( - packed_gen, cfg_text, cfg_img, image_shape, device=device - ) - forward_kwargs = self._forward_kwargs( - packed_gen, cfg_text, cfg_img, gi, gi_cfg_text, gi_cfg_img, params - ) + repeats = len(target) + packed_gen = rl_ops.repeat_context(gen, repeats) + gi, gi_cfg_text, gi_cfg_img = self._build_generation_inputs( + packed_gen, cfg_text, cfg_img, image_shape, device=device + ) + forward_kwargs = self._forward_kwargs( + packed_gen, cfg_text, cfg_img, gi, gi_cfg_text, gi_cfg_img, params + ) + x_t = torch.stack([segment.latents_at(i)[0].to(device) for i in target]) + prev_sample = torch.stack([segment.latents_at(i + 1)[0].to(device) for i in target]) + t_cur = torch.stack([schedule[i] for i in target]) + t_next = torch.stack([schedule[i + 1] for i in target]) + timestep = t_cur[:, None].expand(repeats, x_t.shape[1]).reshape(-1) - x_t = torch.stack([segment.latents_at(i)[0].to(device) for i in chunk]) - prev_sample = torch.stack([segment.latents_at(i + 1)[0].to(device) for i in chunk]) - t_cur = torch.stack([schedule[i] for i in chunk]) - t_next = torch.stack([schedule[i + 1] for i in chunk]) - seq = int(x_t.shape[1]) - timestep = t_cur[:, None].expand(repeats, seq).reshape(-1) - - rl_ops.disable_inference_cache(bagel) - v_t = rl_ops.forward_flow( - bagel, - x_t=x_t.flatten(0, 1), - timestep=timestep, - cfg_text_scale=1.0, - cfg_img_scale=1.0, - **forward_kwargs, - ).view_as(x_t) - _, log_prob, prev_mean = self.strategy.denoise( - noise_pred=v_t, - sample=x_t, - sigma=t_cur, - sigma_next=t_next, - eta=float(params.eta), - prev_sample=prev_sample, - sigma_max=float(sigma_max), - ) - if log_prob is None or prev_mean is None: - raise RuntimeError("Bagel packed replay requires a stochastic FlowGRPO SDE strategy") - log_prob_chunks.append(log_prob.reshape(-1)) - prev_mean_chunks.append(prev_mean) + with self._autocast_ctx(device): + rl_ops.disable_inference_cache(bagel) + v_t = rl_ops.forward_flow( + bagel, + x_t=x_t.flatten(0, 1), + timestep=timestep, + cfg_text_scale=1.0, + cfg_img_scale=1.0, + **forward_kwargs, + ).view_as(x_t) + _, log_prob, prev_mean = self.strategy.denoise( + noise_pred=v_t, + sample=x_t, + sigma=t_cur, + sigma_next=t_next, + eta=float(params.eta), + prev_sample=prev_sample, + sigma_max=float(sigma_max), + ) + if log_prob is None or prev_mean is None: + raise RuntimeError("Bagel packed replay requires a stochastic FlowGRPO SDE strategy") return ReplayResult( - log_probs=torch.cat(log_prob_chunks).unsqueeze(0).to(dtype=self.logprob_dtype), - prev_sample_means=torch.cat(prev_mean_chunks) - .unsqueeze(0) - .to(dtype=self.trajectory_dtype), + log_probs=log_prob.reshape(1, -1).to(dtype=self.logprob_dtype), + prev_sample_means=prev_mean.unsqueeze(0).to(dtype=self.trajectory_dtype), ) gi, gi_cfg_text, gi_cfg_img = self._build_generation_inputs(gen, cfg_text, cfg_img, image_shape, device=device) diff --git a/unirl/models/bagel/pipeline.py b/unirl/models/bagel/pipeline.py index 3872a7aa5..0760a82f2 100644 --- a/unirl/models/bagel/pipeline.py +++ b/unirl/models/bagel/pipeline.py @@ -64,6 +64,7 @@ def __init__( max_prompt_length: int = 8192, cache_t2i_contexts: Optional[bool] = None, context_cache_size: Optional[int] = None, + batch_replay_steps: bool = False, ) -> None: super().__init__() self.bundle = bundle @@ -74,7 +75,7 @@ def __init__( autocast_precision=autocast_precision, trajectory_precision=trajectory_precision, logprob_precision=logprob_precision, - replay_step_pack_size=_cfg_get(getattr(bundle, "config", None), "replay_step_pack_size", 1), + batch_replay_steps=batch_replay_steps, ) self.diffusion = diffusion self.vae_decode = vae_decode if vae_decode is not None else BagelVAEDecodeStage(bundle) From 2fbfdb3a6c0042c300718182a09129be20101a82 Mon Sep 17 00:00:00 2001 From: zhouzhuoxin Date: Wed, 26 Aug 2026 14:46:11 +0800 Subject: [PATCH 3/3] fix(bagel): preserve packed replay sequence boundaries Keep one packed layer traversal while evaluating projection and MLP kernels per sequence, preventing BF16 reduction drift from corrupting replay gradients. --- unirl/models/bagel/diffusion.py | 4 +- .../bagel/vendor/modeling/bagel/bagel.py | 51 +++++++++-- .../vendor/modeling/bagel/qwen2_navit.py | 86 +++++++++++++++---- 3 files changed, 114 insertions(+), 27 deletions(-) diff --git a/unirl/models/bagel/diffusion.py b/unirl/models/bagel/diffusion.py index f9c6bd83d..6565d784a 100644 --- a/unirl/models/bagel/diffusion.py +++ b/unirl/models/bagel/diffusion.py @@ -453,9 +453,7 @@ def replay( gi, gi_cfg_text, gi_cfg_img = self._build_generation_inputs( packed_gen, cfg_text, cfg_img, image_shape, device=device ) - forward_kwargs = self._forward_kwargs( - packed_gen, cfg_text, cfg_img, gi, gi_cfg_text, gi_cfg_img, params - ) + forward_kwargs = self._forward_kwargs(packed_gen, cfg_text, cfg_img, gi, gi_cfg_text, gi_cfg_img, params) x_t = torch.stack([segment.latents_at(i)[0].to(device) for i in target]) prev_sample = torch.stack([segment.latents_at(i + 1)[0].to(device) for i in target]) t_cur = torch.stack([schedule[i] for i in target]) diff --git a/unirl/models/bagel/vendor/modeling/bagel/bagel.py b/unirl/models/bagel/vendor/modeling/bagel/bagel.py index 9d88a8816..3e5e8997f 100644 --- a/unirl/models/bagel/vendor/modeling/bagel/bagel.py +++ b/unirl/models/bagel/vendor/modeling/bagel/bagel.py @@ -24,6 +24,19 @@ from tqdm import tqdm +def _sequence_lengths(lengths) -> List[int]: + if isinstance(lengths, torch.Tensor): + return [int(length) for length in lengths.tolist()] + return [int(length) for length in lengths] + + +def _apply_by_sequence(module, packed_tensor, lengths): + lengths = _sequence_lengths(lengths) + if len(lengths) <= 1: + return module(packed_tensor) + return torch.cat([module(chunk) for chunk in packed_tensor.split(lengths, dim=0)], dim=0) + + class BagelConfig(PretrainedConfig): def __init__( self, @@ -801,9 +814,23 @@ def _forward_flow( raise ValueError( f"timestep must contain one value per latent token, got {tuple(timestep.shape)} for {x_t.shape[0]}" ) - packed_pos_embed = self.latent_pos_embed(packed_vae_position_ids) - packed_timestep_embeds = self.time_embedder(timestep) - x_t = self.vae2llm(x_t) + packed_timestep_embeds + packed_pos_embed + query_lengths = _sequence_lengths(packed_seqlens) + latent_lengths = [length - 2 for length in query_lengths] + packed_pos_embed = _apply_by_sequence( + self.latent_pos_embed, + packed_vae_position_ids, + latent_lengths, + ) + packed_timestep_embeds = _apply_by_sequence( + self.time_embedder, + timestep, + latent_lengths, + ) + x_t = ( + _apply_by_sequence(self.vae2llm, x_t, latent_lengths) + + packed_timestep_embeds + + packed_pos_embed + ) if x_t.dtype != packed_sequence.dtype: x_t = x_t.to(packed_sequence.dtype) packed_sequence[packed_vae_token_indexes] = x_t @@ -832,7 +859,11 @@ def _forward_flow( is_causal=False, **extra_inputs, ) - v_t = self.llm2vae(output.packed_query_sequence) + v_t = _apply_by_sequence( + self.llm2vae, + output.packed_query_sequence, + query_lengths, + ) v_t = v_t[packed_vae_token_indexes] if cfg_text_scale > 1.0: @@ -851,7 +882,11 @@ def _forward_flow( is_causal=False, **extra_inputs, ) - cfg_text_v_t = self.llm2vae(cfg_text_output.packed_query_sequence) + cfg_text_v_t = _apply_by_sequence( + self.llm2vae, + cfg_text_output.packed_query_sequence, + query_lengths, + ) cfg_text_v_t = cfg_text_v_t[packed_vae_token_indexes] if cfg_img_scale > 1.0: @@ -870,7 +905,11 @@ def _forward_flow( is_causal=False, **extra_inputs, ) - cfg_img_v_t = self.llm2vae(cfg_img_output.packed_query_sequence) + cfg_img_v_t = _apply_by_sequence( + self.llm2vae, + cfg_img_output.packed_query_sequence, + query_lengths, + ) cfg_img_v_t = cfg_img_v_t[packed_vae_token_indexes] if cfg_text_scale > 1.0: diff --git a/unirl/models/bagel/vendor/modeling/bagel/qwen2_navit.py b/unirl/models/bagel/vendor/modeling/bagel/qwen2_navit.py index a5d8e0dcc..f1382c391 100644 --- a/unirl/models/bagel/vendor/modeling/bagel/qwen2_navit.py +++ b/unirl/models/bagel/vendor/modeling/bagel/qwen2_navit.py @@ -233,6 +233,36 @@ def pad_sequence(tensor, pad_size): return torch.cat([tensor, pad_tensor], dim=1) +def _sequence_lengths(query_lens) -> List[int]: + if isinstance(query_lens, torch.Tensor): + return [int(length) for length in query_lens.tolist()] + return [int(length) for length in query_lens] + + +def _apply_by_sequence(module, packed_tensor, query_lens): + lengths = _sequence_lengths(query_lens) + if len(lengths) <= 1: + return module(packed_tensor) + return torch.cat([module(chunk) for chunk in packed_tensor.split(lengths, dim=0)], dim=0) + + +def _apply_selected_by_sequence(module, packed_tensor, selected_indexes, query_lens): + lengths = _sequence_lengths(query_lens) + if len(lengths) <= 1: + return module(packed_tensor[selected_indexes]) + + outputs = [] + offset = 0 + for length in lengths: + end = offset + length + mask = (selected_indexes >= offset) & (selected_indexes < end) + indexes = selected_indexes[mask] + if indexes.numel() > 0: + outputs.append(module(packed_tensor[indexes])) + offset = end + return torch.cat(outputs, dim=0) + + class PackedAttention(Qwen2Attention): def __init__(self, config, layer_idx: Optional[int] = None): super().__init__(config, layer_idx) @@ -523,17 +553,26 @@ def forward_inference( packed_key_states = packed_query_sequence.new_zeros((packed_query_sequence.shape[0], self.num_key_value_heads * self.head_dim)) packed_value_states = packed_query_sequence.new_zeros((packed_query_sequence.shape[0], self.num_key_value_heads * self.head_dim)) - packed_text_query_sequence = packed_query_sequence[packed_text_indexes] - packed_vae_query_sequence = packed_query_sequence[packed_vae_token_indexes] - - packed_query_states[packed_text_indexes] = self.q_proj(packed_text_query_sequence) - packed_query_states[packed_vae_token_indexes] = self.q_proj_moe_gen(packed_vae_query_sequence) + packed_query_states[packed_text_indexes] = _apply_selected_by_sequence( + self.q_proj, packed_query_sequence, packed_text_indexes, query_lens + ) + packed_query_states[packed_vae_token_indexes] = _apply_selected_by_sequence( + self.q_proj_moe_gen, packed_query_sequence, packed_vae_token_indexes, query_lens + ) - packed_key_states[packed_text_indexes] = self.k_proj(packed_text_query_sequence) - packed_key_states[packed_vae_token_indexes] = self.k_proj_moe_gen(packed_vae_query_sequence) + packed_key_states[packed_text_indexes] = _apply_selected_by_sequence( + self.k_proj, packed_query_sequence, packed_text_indexes, query_lens + ) + packed_key_states[packed_vae_token_indexes] = _apply_selected_by_sequence( + self.k_proj_moe_gen, packed_query_sequence, packed_vae_token_indexes, query_lens + ) - packed_value_states[packed_text_indexes] = self.v_proj(packed_text_query_sequence) - packed_value_states[packed_vae_token_indexes] = self.v_proj_moe_gen(packed_vae_query_sequence) + packed_value_states[packed_text_indexes] = _apply_selected_by_sequence( + self.v_proj, packed_query_sequence, packed_text_indexes, query_lens + ) + packed_value_states[packed_vae_token_indexes] = _apply_selected_by_sequence( + self.v_proj_moe_gen, packed_query_sequence, packed_vae_token_indexes, query_lens + ) packed_query_states = packed_query_states.view(-1, self.num_heads, self.head_dim) packed_key_states = packed_key_states.view(-1, self.num_key_value_heads, self.head_dim) @@ -599,8 +638,12 @@ def forward_inference( # path is already functional, and the gen input_layernorm/MLP already use this # zeros_like pattern; this mirrors flow_grpo's identical fix. Math is unchanged. packed_attn_output_ = torch.zeros_like(packed_attn_output) - packed_attn_output_[packed_text_indexes] = self.o_proj(packed_attn_output[packed_text_indexes]) - packed_attn_output_[packed_vae_token_indexes] = self.o_proj_moe_gen(packed_attn_output[packed_vae_token_indexes]) + packed_attn_output_[packed_text_indexes] = _apply_selected_by_sequence( + self.o_proj, packed_attn_output, packed_text_indexes, query_lens + ) + packed_attn_output_[packed_vae_token_indexes] = _apply_selected_by_sequence( + self.o_proj_moe_gen, packed_attn_output, packed_vae_token_indexes, query_lens + ) packed_attn_output = packed_attn_output_ if update_past_key_values: @@ -819,14 +862,21 @@ def forward_inference( packed_query_sequence = self.post_attention_layernorm(packed_query_sequence) packed_query_sequence = self.mlp(packed_query_sequence) elif mode == "gen": - packed_text_query_sequence = packed_query_sequence[packed_text_indexes] - packed_vae_query_sequence = packed_query_sequence[packed_vae_token_indexes] - packed_text_query_sequence = self.post_attention_layernorm(packed_text_query_sequence).to(torch.bfloat16) - packed_vae_query_sequence = self.post_attention_layernorm_moe_gen(packed_vae_query_sequence).to(torch.bfloat16) - packed_query_sequence_ = torch.zeros_like(packed_query_sequence).to(torch.bfloat16) - packed_query_sequence_[packed_text_indexes] = self.mlp(packed_text_query_sequence) - packed_query_sequence_[packed_vae_token_indexes] = self.mlp_moe_gen(packed_vae_query_sequence) + packed_query_sequence_[packed_text_indexes] = _apply_selected_by_sequence( + lambda values: self.mlp(self.post_attention_layernorm(values).to(torch.bfloat16)), + packed_query_sequence, + packed_text_indexes, + query_lens, + ) + packed_query_sequence_[packed_vae_token_indexes] = _apply_selected_by_sequence( + lambda values: self.mlp_moe_gen( + self.post_attention_layernorm_moe_gen(values).to(torch.bfloat16) + ), + packed_query_sequence, + packed_vae_token_indexes, + query_lens, + ) packed_query_sequence = packed_query_sequence_ packed_query_sequence = residual + packed_query_sequence