From 7ca7bc8847f4a4ff3b61f5fe94fb48f84f3788db Mon Sep 17 00:00:00 2001 From: CjhHa1 Date: Sat, 25 Jul 2026 18:21:01 +0800 Subject: [PATCH 1/2] perf(bagel): batch trainside UniGRPO rollout forwards Pack independent AR chains and same-shape diffusion samples block-diagonally to amortize live-FSDP all-gathers while preserving the serial path as the default. --- .../bagel_trainside_unigrpo.yaml | 21 +- tests/models/bagel/test_pack_b.py | 360 ++++++++++++++++++ unirl/models/bagel/ar.py | 104 ++++- unirl/models/bagel/diffusion.py | 338 +++++++++++++++- unirl/models/bagel/pipeline.py | 92 +++-- unirl/models/bagel/rl_ops.py | 120 ++++++ 6 files changed, 951 insertions(+), 84 deletions(-) create mode 100644 tests/models/bagel/test_pack_b.py diff --git a/examples/unified_model/bagel_trainside_unigrpo.yaml b/examples/unified_model/bagel_trainside_unigrpo.yaml index a9dfff863..f1018f2a3 100644 --- a/examples/unified_model/bagel_trainside_unigrpo.yaml +++ b/examples/unified_model/bagel_trainside_unigrpo.yaml @@ -26,12 +26,9 @@ # Full fine-tuning: the MoT decoder experts are trained directly (no LoRA). FSDP # shards params + grads + optimizer states across all ranks. # -# NOTE (rebased onto main): the rollout is per-sample navit bs=1 (the branch's -# pack-B / forward_batch_size rollout packing is dropped here — it conflicts with -# main's KV-in-conditions diffusion design; re-add as a follow-up if rollout -# throughput matters). With bs=1 the rollout and replay geometry match, so the -# image side uses old_logp_source=rollout (the on-policy ratio is 1 without a -# pre-update replay). +# Rollout packing is opt-in. forward_batch_size=1 preserves the per-sample path; +# larger values block-diagonally pack text-only thinking chains and same-shape +# diffusion images while conditions continue to hold per-sample KV contexts. # # Launch (multi-node; the launcher sets num_devices = NUM_NODES * GPUS_PER_NODE): # PYTHONPATH=$PWD python -m unirl.train_unified_model \ @@ -83,6 +80,9 @@ pipeline: trajectory_precision: fp32 logprob_precision: fp32 shift: 3.0 + # Pack at most B thinking chains / images into each navit forward. Keep 1 as + # the conservative default; tune per GPU after parity and memory validation. + forward_batch_size: ${oc.env:BAGEL_FWD_BS,1} strategy: _target_: unirl.sde.kernels.FlowSDEStrategy @@ -127,7 +127,7 @@ backend: # transformer (the decoder blocks the bundle unfroze). # Single TRAINSIDE rollout (the M=1 / UniGRPO mode). stage_attrs eval-scopes BOTH -# trainable stages (the same shared transformer). Per-sample navit bs=1. +# trainable stages (the same shared transformer). rollout: _target_: unirl.rollout.engine.trainside.engine.TrainsideRolloutEngine stage_attrs: [diffusion, ar] @@ -167,15 +167,14 @@ algorithm: clip_range: 1.0e-6 # flow trust region (clip range) clip_schedule: constant # Rollout anchor: μ_old (sde_means) and π_old (sde_logp) are the rollout's own - # per-SDE-step recordings. Valid because the rollout is per-sample navit bs=1 — the - # same geometry as the bs=1 train replay — so on-policy μ_old == μ_θ → ratio = 1. - # (Replay anchor would only be needed if rollout ran a different geometry, e.g. pack-B.) + # per-SDE-step recordings. Block-diagonal packing isolates every sequence, so + # each row retains the same replay geometry and on-policy ratio semantics. old_logp_source: rollout mse_weight: 1.5e-5 # velocity-MSE weight (replaces the latent KL) # GRPO-Guard RatioNorm: per-SDE-step normalize the flow ratio (it is otherwise # left-shifted, mean < 1, so clipping never engages). Only bites on the off-policy # update(s), i.e. needs num_updates_per_batch >= 2. μ_old comes from segment.sde_means - # as recorded by the rollout (bs=1, matching the bs=1 replay geometry). + # as recorded by the rollout. ratio_norm: true grad_reweight: false # GRPO-Guard's optional 2nd part (x 1/dt); off by default conditions_cls: diff --git a/tests/models/bagel/test_pack_b.py b/tests/models/bagel/test_pack_b.py new file mode 100644 index 000000000..55566ae98 --- /dev/null +++ b/tests/models/bagel/test_pack_b.py @@ -0,0 +1,360 @@ +import sys +from types import ModuleType, SimpleNamespace + +import torch + +from unirl.models.bagel.ar import BagelARStage, BagelARStep +from unirl.models.bagel.conditions import BagelARConditions, BagelDiffusionConditions +from unirl.models.bagel.diffusion import BagelDiffusionStage, BagelDiffusionStep +from unirl.sde.kernels import FlowSDEStrategy +from unirl.types.sampling import ARSamplingParams + + +class NaiveCache: + def __init__(self, num_layers): + self.key_cache = {index: None for index in range(num_layers)} + self.value_cache = {index: None for index in range(num_layers)} + + @property + def num_layers(self): + return len(self.key_cache) + + +class _FakeLMModel(torch.nn.Module): + def embed_tokens(self, token_ids): + return token_ids.to(torch.float32).unsqueeze(-1) + + +class _FakeLanguageModel(torch.nn.Module): + def __init__(self): + super().__init__() + self.model = _FakeLMModel() + + def forward_inference( + self, + *, + packed_query_sequence, + query_lens, + packed_query_position_ids, + packed_query_indexes, + past_key_values, + key_values_lens, + packed_key_value_indexes, + update_past_key_values, + **_, + ): + batch_size = int(key_values_lens.numel()) + expected_queries = torch.cumsum(key_values_lens, dim=0) + torch.arange(batch_size, dtype=key_values_lens.dtype) + assert torch.equal(packed_query_indexes.cpu(), expected_queries.cpu()) + source_blocks = torch.arange(int(key_values_lens.sum())).split(key_values_lens.tolist()) + expected_keys = torch.cat([block + row for row, block in enumerate(source_blocks)]) + assert torch.equal(packed_key_value_indexes.cpu(), expected_keys.cpu()) + assert torch.equal(query_lens, torch.ones_like(query_lens)) + del packed_query_position_ids + old_tokens = past_key_values.key_cache[0].reshape(-1) + old_blocks = list(old_tokens.split(key_values_lens.tolist())) + query_tokens = packed_query_sequence.reshape(-1) + hidden = torch.stack([query_tokens[row] + old_blocks[row].sum() for row in range(len(old_blocks))]).unsqueeze( + -1 + ) + if update_past_key_values: + merged = torch.cat( + [torch.cat([block, query_tokens[row : row + 1]]) for row, block in enumerate(old_blocks)] + ) + past_key_values.key_cache[0] = merged.reshape(-1, 1, 1) + past_key_values.value_cache[0] = merged.reshape(-1, 1, 1) + return SimpleNamespace(packed_query_sequence=hidden, past_key_values=past_key_values) + + def lm_head(self, hidden): + vocab = 17 + targets = hidden.reshape(-1).long().remainder(vocab) + columns = torch.arange(vocab, dtype=torch.float32, device=hidden.device) + return -(columns.unsqueeze(0) - targets.unsqueeze(1)).abs() + + +class _FakeBagel: + def __init__(self): + self.language_model = _FakeLanguageModel() + self.config = SimpleNamespace(llm_config=SimpleNamespace(num_hidden_layers=1, freeze_und=False)) + + def forward_cache_update_text( + self, + past_key_values, + *, + text_token_lens, + packed_text_ids, + key_values_lens, + **_, + ): + new_blocks = list(packed_text_ids.to(torch.float32).split(text_token_lens.tolist())) + if past_key_values.key_cache[0] is None: + old_blocks = [packed_text_ids.new_zeros(0, dtype=torch.float32) for _ in new_blocks] + else: + old_blocks = list(past_key_values.key_cache[0].reshape(-1).split(key_values_lens.tolist())) + merged = torch.cat([torch.cat([old, new]) for old, new in zip(old_blocks, new_blocks)]) + past_key_values.key_cache[0] = merged.reshape(-1, 1, 1) + past_key_values.value_cache[0] = merged.reshape(-1, 1, 1) + return past_key_values + + +class _FakeBundle: + def __init__(self): + self.model = _FakeBagel() + self.device = torch.device("cpu") + self.new_token_ids = {"bos_token_id": 0, "eos_token_id": 16} + self.transformer = self.model.language_model + + +def _ar_conditions(second_prompt=(7, 8)): + return BagelARConditions( + prompt_splits=[ + [ + {"kind": "text", "ids": torch.tensor([1, 2])}, + {"kind": "text", "ids": torch.tensor([3])}, + ], + [{"kind": "text", "ids": torch.tensor(second_prompt)}], + ] + ) + + +def _ar_stage(forward_batch_size): + bundle = _FakeBundle() + bundle.transformer.eval() + return BagelARStage(model=bundle, forward_batch_size=forward_batch_size) + + +def test_ar_bs1_and_packed_match_tokens_logprobs_and_isolate_samples(): + params = ARSamplingParams( + temperature=0.0, + top_p=1.0, + top_k=0, + max_new_tokens=5, + ) + serial = _ar_stage(1).autoregress(_ar_conditions(), sampling_params=params) + packed = _ar_stage(2).autoregress(_ar_conditions(), sampling_params=params) + + assert torch.equal(serial.tokens, packed.tokens) + assert torch.equal(serial.lengths, packed.lengths) + assert torch.equal(serial.cu_seqlens, packed.cu_seqlens) + torch.testing.assert_close(serial.log_probs, packed.log_probs, rtol=0, atol=0) + + changed_peer = _ar_stage(2).autoregress(_ar_conditions(second_prompt=(12, 13, 14)), sampling_params=params) + first_end = int(packed.cu_seqlens[1]) + assert torch.equal(packed.tokens[:first_end], changed_peer.tokens[:first_end]) + torch.testing.assert_close( + packed.log_probs[:first_end], + changed_peer.log_probs[:first_end], + rtol=0, + atol=0, + ) + + +def test_ar_fbs1_uses_legacy_global_rng_and_packed_uses_generators(monkeypatch): + params = ARSamplingParams( + temperature=0.8, + top_p=0.95, + top_k=8, + max_new_tokens=5, + ) + calls = [] + original_step = BagelARStep.step + + def recording_step(self, logits, *, generators=None): + calls.append(generators) + return original_step(self, logits, generators=generators) + + monkeypatch.setattr(BagelARStep, "step", recording_step) + torch.manual_seed(1234) + _ar_stage(1).autoregress(_ar_conditions(), sampling_params=params) + assert calls and all(generators is None for generators in calls) + + calls.clear() + torch.manual_seed(1234) + _ar_stage(2).autoregress(_ar_conditions(), sampling_params=params) + assert calls and all(generators is not None for generators in calls) + + +def test_ar_vit_input_declines_packing(): + assert BagelARStage._batched_text_ids([[{"kind": "vit", "image": torch.zeros(3, 2, 2)}]]) is None + + +class _FakeStrategy: + def __init__(self): + self.denoise_calls = 0 + + def denoise( + self, + *, + noise_pred, + sample, + prev_sample, + **_, + ): + self.denoise_calls += 1 + mean = sample - 0.25 * noise_pred + output = mean if prev_sample is None else prev_sample + log_prob = -((output - mean) ** 2).mean(dim=(1, 2)) + return output, log_prob, mean + + +def test_diffusion_single_sample_uses_legacy_denoise_path(): + strategy = _FakeStrategy() + step = BagelDiffusionStep() + step.denoise( + strategy, + v_t=torch.ones(3, 4), + x_t=torch.zeros(3, 4), + sigma=torch.tensor(1.0), + sigma_next=torch.tensor(0.5), + sigma_max=torch.tensor(1.0), + eta=0.8, + ) + assert strategy.denoise_calls == 1 + + +def test_diffusion_packed_reduction_matches_bs1_and_isolates_rows(): + step = BagelDiffusionStep() + strategy = _FakeStrategy() + sample = torch.arange(24, dtype=torch.float32).reshape(6, 4) + velocity = torch.arange(24, dtype=torch.float32).flip(0).reshape(6, 4) + kwargs = dict( + sigma=torch.tensor(1.0), + sigma_next=torch.tensor(0.5), + sigma_max=torch.tensor(1.0), + eta=0.8, + ) + + packed = step.denoise( + strategy, + v_t=velocity, + x_t=sample, + n_samples=2, + **kwargs, + ) + serial = [ + step.denoise( + strategy, + v_t=velocity[index * 3 : (index + 1) * 3], + x_t=sample[index * 3 : (index + 1) * 3], + **kwargs, + ) + for index in range(2) + ] + + torch.testing.assert_close(packed[0], torch.cat([value[0] for value in serial])) + torch.testing.assert_close(packed[1], torch.stack([value[1] for value in serial])) + torch.testing.assert_close(packed[2], torch.cat([value[2] for value in serial])) + + changed_velocity = velocity.clone() + changed_velocity[3:] += 1000 + changed = step.denoise( + strategy, + v_t=changed_velocity, + x_t=sample, + n_samples=2, + **kwargs, + ) + torch.testing.assert_close(packed[0][:3], changed[0][:3], rtol=0, atol=0) + torch.testing.assert_close(packed[1][0], changed[1][0], rtol=0, atol=0) + + +def test_diffusion_fixed_per_sequence_rng_matches_serial_steps(monkeypatch): + torch_utils = ModuleType("diffusers.utils.torch_utils") + + def randn_tensor(shape, *, generator, device, dtype): + if isinstance(generator, list): + return torch.cat( + [ + torch.randn((1, *shape[1:]), generator=row_generator, device=device, dtype=dtype) + for row_generator in generator + ] + ) + return torch.randn(shape, generator=generator, device=device, dtype=dtype) + + torch_utils.randn_tensor = randn_tensor + monkeypatch.setitem(sys.modules, "diffusers", ModuleType("diffusers")) + monkeypatch.setitem(sys.modules, "diffusers.utils", ModuleType("diffusers.utils")) + monkeypatch.setitem(sys.modules, "diffusers.utils.torch_utils", torch_utils) + + step = BagelDiffusionStep() + packed_strategy = FlowSDEStrategy() + serial_strategies = [FlowSDEStrategy(), FlowSDEStrategy()] + packed_generators = [torch.Generator().manual_seed(seed) for seed in (11, 29)] + serial_generators = [torch.Generator().manual_seed(seed) for seed in (11, 29)] + packed_sample = torch.arange(24, dtype=torch.float32).reshape(6, 4) / 10 + serial_samples = [packed_sample[:3].clone(), packed_sample[3:].clone()] + schedule = [1.0, 0.8, 0.5, 0.2] + + for current, following in zip(schedule, schedule[1:]): + packed_velocity = packed_sample * 0.1 + 0.25 + packed_sample, packed_logp, _ = step.denoise( + packed_strategy, + v_t=packed_velocity, + x_t=packed_sample, + sigma=torch.tensor(current), + sigma_next=torch.tensor(following), + sigma_max=torch.tensor(0.8), + eta=0.8, + n_samples=2, + generators=packed_generators, + ) + + serial_logps = [] + for index in range(2): + velocity = serial_samples[index] * 0.1 + 0.25 + serial_samples[index], logp, _ = step.denoise( + serial_strategies[index], + v_t=velocity, + x_t=serial_samples[index], + sigma=torch.tensor(current), + sigma_next=torch.tensor(following), + sigma_max=torch.tensor(0.8), + eta=0.8, + generators=[serial_generators[index]], + ) + serial_logps.append(logp) + + torch.testing.assert_close(packed_sample, torch.cat(serial_samples), rtol=0, atol=0) + torch.testing.assert_close(packed_logp, torch.stack(serial_logps), rtol=0, atol=0) + + +def _context(length, start=0): + cache = NaiveCache(1) + if length: + values = torch.arange(start, start + length, dtype=torch.float32).reshape(-1, 1, 1) + cache.key_cache[0] = values + cache.value_cache[0] = values + 100 + return {"kv_lens": [length], "ropes": [length], "past_key_values": cache} + + +def test_diffusion_context_merge_and_pack_fallback_rules(): + merged = BagelDiffusionStage._merge_contexts([_context(2), _context(0), _context(1, 9)]) + assert merged["kv_lens"] == [2, 0, 1] + assert merged["ropes"] == [2, 0, 1] + torch.testing.assert_close( + merged["past_key_values"].key_cache[0].reshape(-1), + torch.tensor([0.0, 1.0, 9.0]), + ) + + same_shape = BagelDiffusionConditions( + gen_contexts=[_context(1), _context(1)], + cfg_text_contexts=[_context(0), _context(0)], + cfg_img_contexts=[_context(1), _context(1)], + prompts=["a", "b"], + image_shapes=[(512, 512), (512, 512)], + ) + mixed_shape = BagelDiffusionConditions( + gen_contexts=same_shape.gen_contexts, + cfg_text_contexts=same_shape.cfg_text_contexts, + cfg_img_contexts=same_shape.cfg_img_contexts, + prompts=same_shape.prompts, + image_shapes=[(512, 512), (384, 640)], + ) + deferred = BagelDiffusionConditions( + prompts=["a", "b"], + image_shapes=[(512, 512), (512, 512)], + ) + + assert BagelDiffusionStage._can_pack_conditions(same_shape) + assert not BagelDiffusionStage._can_pack_conditions(mixed_shape) + assert not BagelDiffusionStage._can_pack_conditions(deferred) diff --git a/unirl/models/bagel/ar.py b/unirl/models/bagel/ar.py index ad7459b45..85f30601f 100644 --- a/unirl/models/bagel/ar.py +++ b/unirl/models/bagel/ar.py @@ -30,6 +30,7 @@ from contextlib import nullcontext from dataclasses import dataclass from dataclasses import field as dc_field +from functools import partial from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple import torch @@ -75,7 +76,12 @@ def __init__(self, *, temperature: float = 1.0, top_p: float = 1.0, top_k: int = self.top_p = float(top_p) self.top_k = int(top_k) - def step(self, logits: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: + def step( + self, + logits: torch.Tensor, + *, + generators: Optional[List[torch.Generator]] = None, + ) -> Tuple[torch.Tensor, torch.Tensor]: if logits.dim() != 2: raise ValueError(f"BagelARStep.step: expected logits shape [B, vocab], got {tuple(logits.shape)}") @@ -107,7 +113,18 @@ def step(self, logits: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: scaled = torch.full_like(scaled, float("-inf")).scatter(-1, sorted_idx, sorted_vals) probs = F.softmax(scaled, dim=-1) - token_id = torch.multinomial(probs, num_samples=1).squeeze(-1) + if generators is None: + token_id = torch.multinomial(probs, num_samples=1).squeeze(-1) + else: + if len(generators) != int(probs.shape[0]): + raise ValueError(f"BagelARStep.step: got {len(generators)} generators for batch {int(probs.shape[0])}.") + token_id = torch.cat( + [ + torch.multinomial(probs[row], num_samples=1, generator=generator) + for row, generator in enumerate(generators) + ], + dim=0, + ) log_prob = log_probs_full.gather(-1, token_id.unsqueeze(-1)).squeeze(-1) return token_id, log_prob @@ -128,10 +145,17 @@ def __init__( autocast_precision: str = "bf16", logprob_precision: str = "fp32", replay_mode: str = "train", + forward_batch_size: Optional[int] = 1, ) -> None: self.model = model self.autocast_dtype = parse_torch_dtype(autocast_precision, field_name="BagelARStage.autocast_precision") self.logprob_dtype = parse_torch_dtype(logprob_precision, field_name="BagelARStage.logprob_precision") + forward_batch_size = 1 if forward_batch_size is None else int(forward_batch_size) + require( + forward_batch_size >= 1, + f"BagelARStage.forward_batch_size must be >= 1; got {forward_batch_size!r}.", + ) + self.forward_batch_size = forward_batch_size # Replay scorer for the GRPO ratio's new_logp: # "train" — one grad forward_train per sample (nested mask: image full + # text causal); the und path INCLUDING the image is trained. @@ -199,6 +223,21 @@ def _resolve_stop_ids(self, params: Optional[BagelARParams], sampling_params: AR ids.append(int(self.model.new_token_ids["eos_token_id"])) # <|im_end|>, as in the vendored gen_text return list(dict.fromkeys(ids)) + @staticmethod + def _batched_text_ids(prompt_splits: List[List[Dict[str, Any]]]) -> Optional[List[torch.Tensor]]: + """Concatenate each sample's text splits, or decline packing for ViT/empty inputs.""" + out: List[torch.Tensor] = [] + for splits in prompt_splits: + ids: List[torch.Tensor] = [] + for split in splits: + if split.get("kind") != "text": + return None + ids.append(split["ids"].reshape(-1).to(dtype=torch.long)) + if not ids: + return None + out.append(torch.cat(ids, dim=0)) + return out + # ------------------------------------------------------------------ # Rollout # ------------------------------------------------------------------ @@ -225,22 +264,57 @@ def autoregress( stop_ids = self._resolve_stop_ids(params, sampling_params) start_id = int(self.model.new_token_ids["bos_token_id"]) + text_id_lists = self._batched_text_ids(conditions.prompt_splits) + use_batched = self.forward_batch_size > 1 and text_id_lists is not None and len(text_id_lists) > 1 + generated: List[List[int]] = [] logps: List[List[float]] = [] with torch.no_grad(), self._autocast_ctx(device): - for splits in conditions.prompt_splits: - ctx = self._prefill(splits, device=device) - tokens_i, logps_i = rl_ops.decode_text( - bagel, - ctx, - start_token_id=start_id, - sample_fn=step.step, - max_new_tokens=int(sampling_params.max_new_tokens), - stop_ids=stop_ids, - device=device, - ) - generated.append(tokens_i) - logps.append(logps_i) + if use_batched: + for start in range(0, len(text_id_lists), self.forward_batch_size): + id_chunk = text_id_lists[start : start + self.forward_batch_size] + generator_chunk: Optional[List[torch.Generator]] = None + if step.temperature > 0.0 and len(id_chunk) > 1: + seeds = torch.randint( + 0, + (1 << 63) - 1, + (len(id_chunk),), + dtype=torch.int64, + device=device, + ).cpu() + generator_chunk = [] + for seed in seeds.tolist(): + generator = torch.Generator(device=device) + generator.manual_seed(int(seed)) + generator_chunk.append(generator) + ctx = rl_ops.prefill_text_batched(bagel, id_chunk, device=device) + tokens, token_logps = rl_ops.decode_text_batched( + bagel, + ctx, + start_token_id=start_id, + sample_fn=( + partial(step.step, generators=generator_chunk) if generator_chunk is not None else step.step + ), + max_new_tokens=int(sampling_params.max_new_tokens), + stop_ids=stop_ids, + device=device, + ) + generated.extend(tokens) + logps.extend(token_logps) + else: + for splits in conditions.prompt_splits: + ctx = self._prefill(splits, device=device) + tokens_i, logps_i = rl_ops.decode_text( + bagel, + ctx, + start_token_id=start_id, + sample_fn=step.step, + max_new_tokens=int(sampling_params.max_new_tokens), + stop_ids=stop_ids, + device=device, + ) + generated.append(tokens_i) + logps.append(logps_i) return TextSegment.pack( tokens=[torch.tensor(t, dtype=torch.long, device=device) for t in generated], diff --git a/unirl/models/bagel/diffusion.py b/unirl/models/bagel/diffusion.py index d753375ed..c74688b1a 100644 --- a/unirl/models/bagel/diffusion.py +++ b/unirl/models/bagel/diffusion.py @@ -172,6 +172,8 @@ def denoise( sigma_max: torch.Tensor, eta: float, prev_sample: Optional[torch.Tensor] = None, + n_samples: int = 1, + generators: Optional[List[torch.Generator]] = None, ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[torch.Tensor]]: """One SDE transition via the shared ``strategy.denoise`` over packed latents. @@ -184,19 +186,67 @@ def denoise( ``log_prob`` / ``prev_sample_mean`` are ``None`` for deterministic (``eta < 1e-7``) steps. """ - prev, log_prob, prev_mean = strategy.denoise( - noise_pred=v_t.unsqueeze(0), - sample=x_t.unsqueeze(0), - sigma=sigma, - sigma_next=sigma_next, - eta=float(eta), - prev_sample=None if prev_sample is None else prev_sample.unsqueeze(0), - sigma_max=float(sigma_max), + require(n_samples >= 1, f"BagelDiffusionStep.denoise: n_samples must be >= 1; got {n_samples}.") + require( + int(x_t.shape[0]) % int(n_samples) == 0, + f"BagelDiffusionStep.denoise: packed token count {int(x_t.shape[0])} " + f"is not divisible by n_samples={n_samples}.", ) + seq = int(x_t.shape[0]) // int(n_samples) + channels = int(x_t.shape[-1]) + sample = x_t.reshape(int(n_samples), seq, channels) + noise_pred = v_t.reshape(int(n_samples), seq, channels) + replay_sample = None if prev_sample is None else prev_sample.reshape(int(n_samples), seq, channels) + if generators is None or replay_sample is not None: + prev, log_prob, prev_mean = strategy.denoise( + noise_pred=noise_pred, + sample=sample, + sigma=sigma, + sigma_next=sigma_next, + eta=float(eta), + prev_sample=replay_sample, + sigma_max=float(sigma_max), + ) + else: + if len(generators) != int(n_samples): + raise ValueError( + f"BagelDiffusionStep.denoise: got {len(generators)} generators for n_samples={n_samples}." + ) + input_dtype = sample.dtype + noise_pred_f32 = noise_pred.float() + sample_f32 = sample.float() + sigma_f32 = sigma.float().reshape(1) + sigma_next_f32 = sigma_next.float().reshape(1) + while sigma_f32.dim() < sample_f32.dim(): + sigma_f32 = sigma_f32.unsqueeze(-1) + sigma_next_f32 = sigma_next_f32.unsqueeze(-1) + prev, prev_mean, std_var = strategy.step( + noise_pred=noise_pred_f32, + sample=sample_f32, + sigma=sigma_f32, + sigma_next=sigma_next_f32, + eta=float(eta), + prev_sample=None, + generator=generators, + sigma_max=float(sigma_max), + ) + prev, log_prob = strategy._finalize_logp( + prev_sample=prev, + prev_sample_mean=prev_mean, + std_var=std_var, + eta=float(eta), + input_dtype=input_dtype, + ) + if n_samples == 1: + return ( + prev.reshape(seq, channels), + None if log_prob is None else log_prob.reshape(()), + None if prev_mean is None else prev_mean.reshape(seq, channels), + ) return ( - prev.squeeze(0), - None if log_prob is None else log_prob.reshape(()), - None if prev_mean is None else prev_mean.squeeze(0), + prev.reshape(int(n_samples) * seq, channels), + None if log_prob is None else log_prob.reshape(int(n_samples)), + None if prev_mean is None else prev_mean.reshape(int(n_samples) * seq, channels), ) def step_with_logp( @@ -213,6 +263,8 @@ def step_with_logp( cfg_text_scale: float, cfg_img_scale: float, forward_kwargs: Dict[str, Any], + n_samples: int = 1, + generators: Optional[List[torch.Generator]] = None, ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[torch.Tensor]]: """Run ``predict_velocity`` then ``denoise`` for one step. @@ -237,6 +289,8 @@ def step_with_logp( sigma_max=sigma_max, eta=eta, prev_sample=prev_sample, + n_samples=n_samples, + generators=generators, ) @@ -285,6 +339,25 @@ def _autocast_ctx(self, device: torch.device): return torch.autocast("cuda", self.autocast_dtype) return nullcontext() + @staticmethod + def _sampling_generators(device: torch.device, count: int) -> List[torch.Generator]: + """Fork stable per-sequence RNG streams from the current device RNG.""" + generators: List[torch.Generator] = [] + for _ in range(int(count)): + seed = int( + torch.randint( + 0, + (1 << 63) - 1, + (), + dtype=torch.int64, + device=device, + ).item() + ) + generator = torch.Generator(device=device) + generator.manual_seed(seed) + generators.append(generator) + return generators + def _build_contexts_from_prompt(self, prompt: str) -> Tuple[Any, Any, Any]: """Rebuild the three KV contexts (gen / cfg_text / cfg_img) from a prompt. @@ -443,6 +516,21 @@ def diffuse( indices : [K] stored frame step indices sigmas : [T+1] the full schedule """ + if conditions.batch_size > 1: + if self._can_pack_conditions(conditions): + return self._diffuse_batched( + conditions, + schedule=schedule, + params=params, + initial_latents=initial_latents, + ) + return self._diffuse_serial_batch( + conditions, + schedule=schedule, + params=params, + initial_latents=initial_latents, + ) + bagel = self.model.model device = torch.device(self.model.device) schedule = schedule.to(device) @@ -526,6 +614,234 @@ def diffuse( sde_indices=sde_indices, ) + # ------------------------------------------------------------------ + # Block-diagonal rollout batching + # ------------------------------------------------------------------ + + @staticmethod + def _can_pack_conditions(conditions: BagelDiffusionConditions) -> bool: + """Only opaque, same-shape contexts have an unambiguous packed latent geometry.""" + if not conditions.has_contexts() or conditions.batch_size < 2: + return False + if len(conditions.gen_contexts) != conditions.batch_size: + return False + shapes = [tuple(shape) for shape in conditions.image_shapes] + return len(shapes) == conditions.batch_size and len(set(shapes)) == 1 + + @staticmethod + def _stack_segments(segments: List[LatentSegment]) -> LatentSegment: + if len(segments) == 1: + return segments[0] + return LatentSegment( + latents=torch.cat([segment.latents for segment in segments], dim=0), + sigmas=segments[0].sigmas, + indices=segments[0].indices, + sde_logp=( + torch.cat([segment.sde_logp for segment in segments], dim=0) + if segments[0].sde_logp is not None + else None + ), + sde_means=( + torch.cat([segment.sde_means for segment in segments], dim=0) + if segments[0].sde_means is not None + else None + ), + sde_indices=segments[0].sde_indices, + ) + + def _diffuse_serial_batch( + self, + conditions: BagelDiffusionConditions, + *, + schedule: torch.Tensor, + params: BagelDiffusionParams, + initial_latents: Optional[torch.Tensor], + ) -> LatentSegment: + """Fallback for deferred contexts or mixed image shapes.""" + segments: List[LatentSegment] = [] + for index in range(conditions.batch_size): + if conditions.has_contexts(): + gen = conditions.gen_contexts[index] + cfg_text = ( + conditions.cfg_text_contexts[index] + if conditions.cfg_text_contexts and conditions.cfg_text_contexts[index] is not None + else gen + ) + cfg_img = ( + conditions.cfg_img_contexts[index] + if conditions.cfg_img_contexts and conditions.cfg_img_contexts[index] is not None + else gen + ) + condition = BagelDiffusionConditions.for_sample( + gen_context=gen, + cfg_text_context=cfg_text, + cfg_img_context=cfg_img, + prompt=conditions.prompts[index] if conditions.prompts else None, + image_shape=tuple(conditions.image_shapes[index]), + ) + else: + condition = BagelDiffusionConditions( + prompts=[conditions.prompts[index]], + image_shapes=[tuple(conditions.image_shapes[index])], + ) + initial = initial_latents[index] if initial_latents is not None else None + segments.append(self.diffuse(condition, schedule=schedule, params=params, initial_latents=initial)) + return self._stack_segments(segments) + + @staticmethod + def _merge_contexts(contexts: List[Any]) -> Dict[str, Any]: + """Merge per-sample NaiveCache objects while preserving empty CFG branches.""" + kv_lens = [int(context["kv_lens"][0]) for context in contexts] + ropes = [int(context["ropes"][0]) for context in contexts] + caches = [context["past_key_values"] for context in contexts] + merged = type(caches[0])(int(caches[0].num_layers)) + for layer in range(int(caches[0].num_layers)): + keys = [cache.key_cache[layer] for cache in caches if cache.key_cache[layer] is not None] + values = [cache.value_cache[layer] for cache in caches if cache.value_cache[layer] is not None] + merged.key_cache[layer] = torch.cat(keys, dim=0) if keys else None + merged.value_cache[layer] = torch.cat(values, dim=0) if values else None + return {"kv_lens": kv_lens, "ropes": ropes, "past_key_values": merged} + + def _build_generation_inputs_batched( + self, + gen: Any, + cfg_text: Any, + cfg_img: Any, + image_shapes: List[Tuple[int, int]], + *, + device: torch.device, + ) -> Tuple[Dict[str, Any], Dict[str, Any], Dict[str, Any]]: + """Build the vendor's packed latent indexes for a multi-sequence context.""" + bagel = self.model.model + gi = bagel.prepare_vae_latent( + curr_kvlens=gen["kv_lens"], + curr_rope=gen["ropes"], + image_sizes=image_shapes, + 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_shapes, + ) + gi_cfg_img = bagel.prepare_vae_latent_cfg( + curr_kvlens=cfg_img["kv_lens"], + curr_rope=cfg_img["ropes"], + image_sizes=image_shapes, + ) + return _to_device(gi, device), _to_device(gi_cfg_text, device), _to_device(gi_cfg_img, device) + + def _diffuse_batched( + self, + conditions: BagelDiffusionConditions, + *, + schedule: torch.Tensor, + params: BagelDiffusionParams, + initial_latents: Optional[torch.Tensor] = None, + ) -> LatentSegment: + """Run one block-diagonal navit forward per step for same-shape images.""" + bagel = self.model.model + device = torch.device(self.model.device) + schedule = schedule.to(device) + num_steps = int(schedule.shape[0]) - 1 + require( + num_steps == int(params.num_inference_steps), + f"BagelDiffusionStage._diffuse_batched: schedule length {schedule.shape[0]} != " + f"num_inference_steps+1 ({int(params.num_inference_steps) + 1})", + ) + sigma_max = schedule[1] if int(schedule.shape[0]) > 1 else schedule[0] + sde_set = {int(index) for index in (params.sde_indices or [])} + sde_sorted = sorted(sde_set) + batch_size = int(conditions.batch_size) + + gen_contexts = list(conditions.gen_contexts) + cfg_text_contexts = [ + conditions.cfg_text_contexts[index] + if conditions.cfg_text_contexts and conditions.cfg_text_contexts[index] is not None + else gen_contexts[index] + for index in range(batch_size) + ] + cfg_img_contexts = [ + conditions.cfg_img_contexts[index] + if conditions.cfg_img_contexts and conditions.cfg_img_contexts[index] is not None + else gen_contexts[index] + for index in range(batch_size) + ] + image_shapes = [tuple(shape) for shape in conditions.image_shapes] + gen = self._merge_contexts(gen_contexts) + cfg_text = self._merge_contexts(cfg_text_contexts) + cfg_img = self._merge_contexts(cfg_img_contexts) + gi, gi_cfg_text, gi_cfg_img = self._build_generation_inputs_batched( + gen, + cfg_text, + cfg_img, + image_shapes, + device=device, + ) + forward_kwargs = self._forward_kwargs(gen, cfg_text, cfg_img, gi, gi_cfg_text, gi_cfg_img, params) + + if initial_latents is None: + x_t = gi["packed_init_noises"].to(device=device, dtype=self.trajectory_dtype) + else: + x_t = initial_latents.to(device=device, dtype=self.trajectory_dtype).reshape( + -1, int(initial_latents.shape[-1]) + ) + channels = int(x_t.shape[-1]) + require( + int(x_t.shape[0]) % batch_size == 0, + "BagelDiffusionStage._diffuse_batched: packed latent tokens must divide evenly by batch size.", + ) + seq = int(x_t.shape[0]) // batch_size + + sampling_generators = self._sampling_generators(device, batch_size) + self.strategy.init_schedule(schedule) + needed = set(compute_trajectory_positions(sde_set, num_steps)) + needed.add(num_steps) + stored_pairs: List[Tuple[int, torch.Tensor]] = [] + if 0 in needed: + stored_pairs.append((0, x_t.detach().clone().reshape(batch_size, seq, channels))) + sde_logps: List[torch.Tensor] = [] + sde_means: List[torch.Tensor] = [] + + with torch.no_grad(), self._autocast_ctx(device): + for index in range(num_steps): + t_cur = schedule[index] + t_next = schedule[index + 1] + cfg_text_scale, cfg_img_scale = self._gated_cfg_scales(float(t_cur.item()), params) + x_t, log_prob, prev_mean = self.step.step_with_logp( + bagel, + self.strategy, + x_t=x_t, + prev_sample=None, + t_cur=t_cur, + t_next=t_next, + sigma_max=sigma_max, + eta=float(params.eta) if index in sde_set else 0.0, + cfg_text_scale=cfg_text_scale, + cfg_img_scale=cfg_img_scale, + forward_kwargs=forward_kwargs, + n_samples=batch_size, + generators=sampling_generators, + ) + x_t = x_t.to(dtype=self.trajectory_dtype) + if (index + 1) in needed: + stored_pairs.append((index + 1, x_t.detach().clone().reshape(batch_size, seq, channels))) + if log_prob is not None: + sde_logps.append(log_prob.to(dtype=self.logprob_dtype)) + if prev_mean is not None: + sde_means.append( + prev_mean.detach().reshape(batch_size, seq, channels).to(dtype=self.trajectory_dtype) + ) + + return LatentSegment( + latents=torch.stack([tensor for _, tensor in stored_pairs], dim=1), + sigmas=schedule, + indices=torch.tensor([index for index, _ in stored_pairs], dtype=torch.long, device=device), + sde_logp=torch.stack(sde_logps, dim=1) if sde_logps else None, + sde_means=torch.stack(sde_means, dim=1) if sde_means else None, + sde_indices=(torch.tensor(sde_sorted, dtype=torch.long, device=device) if sde_sorted else None), + ) + # ------------------------------------------------------------------ # Replay # ------------------------------------------------------------------ diff --git a/unirl/models/bagel/pipeline.py b/unirl/models/bagel/pipeline.py index 5918fd0a4..ebd7e4730 100644 --- a/unirl/models/bagel/pipeline.py +++ b/unirl/models/bagel/pipeline.py @@ -104,8 +104,13 @@ def __init__( replay_mode: str = "train", cache_t2i_contexts: Optional[bool] = None, context_cache_size: Optional[int] = None, + forward_batch_size: Optional[int] = 1, ) -> None: super().__init__() + forward_batch_size = 1 if forward_batch_size is None else int(forward_batch_size) + if forward_batch_size < 1: + raise ValueError(f"BagelPipeline.forward_batch_size must be >= 1; got {forward_batch_size!r}.") + self.forward_batch_size = forward_batch_size self.bundle = bundle if diffusion is None: diffusion = BagelDiffusionStage( @@ -127,6 +132,7 @@ def __init__( autocast_precision=autocast_precision, logprob_precision=logprob_precision, replay_mode=replay_mode, + forward_batch_size=self.forward_batch_size, ) self.autocast_precision = autocast_precision # FlowMatch time-shift for the σ schedule policy (read by the hosting engine @@ -510,29 +516,45 @@ def _diffuse_and_decode( """ device = torch.device(self.bundle.device) schedule = req.sigmas.to(device) - initial = NoiseRecipe.from_rollout_req(req).resolve(device=device, dtype=torch.float32) + initial = NoiseRecipe.from_rollout_req(req).for_batch(len(contexts)).resolve(device=device, dtype=torch.float32) + if initial is not None and int(initial.shape[0]) != len(contexts): + if len(contexts) % int(initial.shape[0]) != 0: + raise ValueError( + "BagelPipeline._diffuse_and_decode: initial latent batch " + f"{int(initial.shape[0])} cannot align to {len(contexts)} image samples." + ) + initial = initial.repeat_interleave(len(contexts) // int(initial.shape[0]), dim=0) gen_list: List[Any] = [] cfg_text_list: List[Any] = [] cfg_img_list: List[Any] = [] shapes: List[Tuple[int, int]] = [] - segments: List[LatentSegment] = [] - for i, (gen_ctx, cfg_text_ctx, cfg_img_ctx) in enumerate(contexts): - cond_i = BagelDiffusionConditions.for_sample( - gen_context=gen_ctx, - cfg_text_context=cfg_text_ctx, - cfg_img_context=cfg_img_ctx, - image_shape=image_shape, - prompt=prompts[i], - ) - x0_i = initial[i] if initial is not None else None - seg_i = self.diffusion.diffuse(cond_i, schedule=schedule, params=params, initial_latents=x0_i) - segments.append(seg_i) + for gen_ctx, cfg_text_ctx, cfg_img_ctx in contexts: gen_list.append(gen_ctx) cfg_text_list.append(cfg_text_ctx) cfg_img_list.append(cfg_img_ctx) shapes.append(image_shape) + segments: List[LatentSegment] = [] + for start in range(0, len(contexts), self.forward_batch_size): + end = min(start + self.forward_batch_size, len(contexts)) + cond_chunk = BagelDiffusionConditions( + gen_contexts=gen_list[start:end], + cfg_text_contexts=cfg_text_list[start:end], + cfg_img_contexts=cfg_img_list[start:end], + prompts=list(prompts[start:end]), + image_shapes=shapes[start:end], + ) + initial_chunk = initial[start:end] if initial is not None else None + segments.append( + self.diffusion.diffuse( + cond_chunk, + schedule=schedule, + params=params, + initial_latents=initial_chunk, + ) + ) + segment = self._batch_segments(segments) conditions = BagelDiffusionConditions( gen_contexts=gen_list, @@ -789,8 +811,8 @@ class BagelUniPipeline(BagelPipeline): (:meth:`_build_think_contexts` for the gen/cfg KV contexts, :meth:`_detokenize`, ``self.ar`` / ``self.diffusion`` / ``self.vae_decode``, :meth:`_batch_segments`, :meth:`build_schedule_policy`); only the prompt-level N×M fan-out + lineage is - layered on top of the single-sample ``_generate_t2ti``. Per-sample navit ``bs=1`` - (no pack-B); each image draws its own x_T inside ``diffuse``. + layered on top of the single-sample ``_generate_t2ti``. ``forward_batch_size`` + optionally packs text-only thinking chains and same-shape images block-diagonally. """ def generate(self, req: RolloutReq) -> RolloutResp: @@ -859,40 +881,16 @@ def generate(self, req: RolloutReq) -> RolloutResp: img_prompts = [prompts[i // n_rewrites] for i in range(n_ar) for _ in range(n_images)] img_thinks = [thinking.texts[i] for i in range(n_ar) for _ in range(n_images)] - device = torch.device(self.bundle.device) - schedule = req.sigmas.to(device) - - # Diffuse per image (navit bs=1; each draws its own x_T inside diffuse) over the - # native think contexts, then batch the per-sample segments into the image track. - gen_list: List[Any] = [] - cfg_text_list: List[Any] = [] - cfg_img_list: List[Any] = [] - shapes: List[Tuple[int, int]] = [] - segments: List[LatentSegment] = [] + contexts: List[Tuple[Any, Any, Any]] = [] for prompt, think in zip(img_prompts, img_thinks): - gen_ctx, cfg_text_ctx, cfg_img_ctx = self._build_think_contexts(GEN_THINK_SYSTEM_PROMPT, prompt, think) - cond_i = BagelDiffusionConditions.for_sample( - gen_context=gen_ctx, - cfg_text_context=cfg_text_ctx, - cfg_img_context=cfg_img_ctx, - image_shape=image_shape, - prompt=prompt, - ) - segments.append(self.diffusion.diffuse(cond_i, schedule=schedule, params=diff_params, initial_latents=None)) - gen_list.append(gen_ctx) - cfg_text_list.append(cfg_text_ctx) - cfg_img_list.append(cfg_img_ctx) - shapes.append(image_shape) - - segment = self._batch_segments(segments) - conditions = BagelDiffusionConditions( - gen_contexts=gen_list, - cfg_text_contexts=cfg_text_list, - cfg_img_contexts=cfg_img_list, - prompts=list(img_prompts), - image_shapes=shapes, + contexts.append(self._build_think_contexts(GEN_THINK_SYSTEM_PROMPT, prompt, think)) + segment, conditions, images = self._diffuse_and_decode( + contexts, + prompts=img_prompts, + params=diff_params, + req=req, + image_shape=image_shape, ) - images = self.vae_decode.decode(segment, image_shape=image_shape) image_track = _track_with_field(img_shell, "segment", segment) image_track = _track_with_field(image_track, "decoded", images) diff --git a/unirl/models/bagel/rl_ops.py b/unirl/models/bagel/rl_ops.py index 9eaa6b6b6..dd3c92c6a 100644 --- a/unirl/models/bagel/rl_ops.py +++ b/unirl/models/bagel/rl_ops.py @@ -71,10 +71,12 @@ __all__ = [ "decode_text", + "decode_text_batched", "disable_inference_cache", "forward_flow", "init_und_context", "pack_und_forward_inputs", + "prefill_text_batched", "prefill_text_split", "prefill_vit_split", "require_inference_dispatch", @@ -337,6 +339,124 @@ def decode_text( return tokens, logps +def prefill_text_batched( + model: Any, + id_lists: List[torch.Tensor], + *, + device: torch.device, +) -> Dict[str, Any]: + """Prefill text-only prompts into one block-diagonal KV context.""" + if not id_lists: + raise ValueError("prefill_text_batched: id_lists must be non-empty.") + + packed_text_ids: List[int] = [] + packed_position_ids: List[int] = [] + packed_text_indexes: List[int] = [] + text_token_lens: List[int] = [] + kv_lens: List[int] = [] + base = 0 + for ids in id_lists: + flat = torch.as_tensor(ids, dtype=torch.long).reshape(-1).tolist() + length = len(flat) + if length == 0: + raise ValueError("prefill_text_batched: empty prompt id list.") + packed_text_ids.extend(int(token) for token in flat) + packed_position_ids.extend(range(length)) + packed_text_indexes.extend(range(base, base + length)) + text_token_lens.append(length) + kv_lens.append(length) + base += length + + inputs = _to_device( + { + "text_token_lens": torch.tensor(text_token_lens, dtype=torch.int), + "packed_text_ids": torch.tensor(packed_text_ids, dtype=torch.long), + "packed_text_position_ids": torch.tensor(packed_position_ids, dtype=torch.long), + "packed_text_indexes": torch.tensor(packed_text_indexes, dtype=torch.long), + "packed_key_value_indexes": torch.zeros(0, dtype=torch.long), + "key_values_lens": torch.zeros(len(id_lists), dtype=torch.int), + }, + device, + ) + fresh_cache = init_und_context(model)["past_key_values"] + past = _raw(type(model).forward_cache_update_text)(model, fresh_cache, **inputs) + return {"kv_lens": kv_lens, "ropes": list(kv_lens), "past_key_values": past} + + +def decode_text_batched( + model: Any, + ctx: Dict[str, Any], + *, + start_token_id: int, + sample_fn: Callable[[torch.Tensor], Tuple[torch.Tensor, torch.Tensor]], + max_new_tokens: int, + stop_ids: List[int], + device: torch.device, +) -> Tuple[List[List[int]], List[List[float]]]: + """Decode a text-only batch with block-diagonal KV indexing.""" + require_inference_dispatch(model) + disable_inference_cache(model) + lm = model.language_model + batch_size = len(ctx["kv_lens"]) + if batch_size < 1: + raise ValueError("decode_text_batched: empty context batch.") + + kv_lens = torch.tensor(ctx["kv_lens"], dtype=torch.int, device=device) + positions = torch.tensor(ctx["ropes"], dtype=torch.long, device=device) + packed_kv_indexes = torch.arange(int(kv_lens.sum().item()), dtype=torch.long, device=device) + past = ctx["past_key_values"] + stop_set = {int(token) for token in stop_ids} + + current = torch.full((batch_size,), int(start_token_id), dtype=torch.long, device=device) + tokens: List[List[int]] = [[] for _ in range(batch_size)] + logps: List[List[float]] = [[] for _ in range(batch_size)] + done = [False] * batch_size + + for _ in range(int(max_new_tokens)): + embeddings = lm.model.embed_tokens(current) + query_lens = torch.ones(batch_size, dtype=torch.int, device=device) + query_indexes = torch.cumsum(kv_lens, dim=0) + torch.arange(batch_size, dtype=kv_lens.dtype, device=device) + + blocks = list(packed_kv_indexes.split(kv_lens.tolist(), dim=0)) + shifted_blocks = [block + i for i, block in enumerate(blocks)] + shifted_kv_indexes = torch.cat(shifted_blocks, dim=0) + out = lm.forward_inference( + packed_query_sequence=embeddings, + query_lens=query_lens, + packed_query_position_ids=positions, + packed_query_indexes=query_indexes, + past_key_values=past, + key_values_lens=kv_lens, + packed_key_value_indexes=shifted_kv_indexes, + update_past_key_values=True, + is_causal=True, + mode="und", + ) + past = out.past_key_values + logits = lm.lm_head(out.packed_query_sequence) + token_ids, token_logps = sample_fn(logits) + + for row in range(batch_size): + if done[row]: + continue + token = int(token_ids[row].item()) + tokens[row].append(token) + logps[row].append(float(token_logps[row].item())) + if token in stop_set: + done[row] = True + + current = token_ids.to(device=device, dtype=torch.long).reshape(batch_size) + old_blocks = list(shifted_kv_indexes.split(kv_lens.tolist(), dim=0)) + packed_kv_indexes = torch.cat( + [torch.cat([block, block[-1:] + 1], dim=0) for block in old_blocks], + dim=0, + ) + kv_lens = kv_lens + 1 + positions = positions + 1 + + return tokens, logps + + def score_response( model: Any, ctx: Dict[str, Any], From 347a22e9865f91a8a7f68e5645955604c44f0eb2 Mon Sep 17 00:00:00 2001 From: leviking98z-rgb Date: Sun, 26 Jul 2026 20:11:05 +0800 Subject: [PATCH 2/2] chore: remove PR test file --- tests/models/bagel/test_pack_b.py | 360 ------------------------------ 1 file changed, 360 deletions(-) delete mode 100644 tests/models/bagel/test_pack_b.py diff --git a/tests/models/bagel/test_pack_b.py b/tests/models/bagel/test_pack_b.py deleted file mode 100644 index 55566ae98..000000000 --- a/tests/models/bagel/test_pack_b.py +++ /dev/null @@ -1,360 +0,0 @@ -import sys -from types import ModuleType, SimpleNamespace - -import torch - -from unirl.models.bagel.ar import BagelARStage, BagelARStep -from unirl.models.bagel.conditions import BagelARConditions, BagelDiffusionConditions -from unirl.models.bagel.diffusion import BagelDiffusionStage, BagelDiffusionStep -from unirl.sde.kernels import FlowSDEStrategy -from unirl.types.sampling import ARSamplingParams - - -class NaiveCache: - def __init__(self, num_layers): - self.key_cache = {index: None for index in range(num_layers)} - self.value_cache = {index: None for index in range(num_layers)} - - @property - def num_layers(self): - return len(self.key_cache) - - -class _FakeLMModel(torch.nn.Module): - def embed_tokens(self, token_ids): - return token_ids.to(torch.float32).unsqueeze(-1) - - -class _FakeLanguageModel(torch.nn.Module): - def __init__(self): - super().__init__() - self.model = _FakeLMModel() - - def forward_inference( - self, - *, - packed_query_sequence, - query_lens, - packed_query_position_ids, - packed_query_indexes, - past_key_values, - key_values_lens, - packed_key_value_indexes, - update_past_key_values, - **_, - ): - batch_size = int(key_values_lens.numel()) - expected_queries = torch.cumsum(key_values_lens, dim=0) + torch.arange(batch_size, dtype=key_values_lens.dtype) - assert torch.equal(packed_query_indexes.cpu(), expected_queries.cpu()) - source_blocks = torch.arange(int(key_values_lens.sum())).split(key_values_lens.tolist()) - expected_keys = torch.cat([block + row for row, block in enumerate(source_blocks)]) - assert torch.equal(packed_key_value_indexes.cpu(), expected_keys.cpu()) - assert torch.equal(query_lens, torch.ones_like(query_lens)) - del packed_query_position_ids - old_tokens = past_key_values.key_cache[0].reshape(-1) - old_blocks = list(old_tokens.split(key_values_lens.tolist())) - query_tokens = packed_query_sequence.reshape(-1) - hidden = torch.stack([query_tokens[row] + old_blocks[row].sum() for row in range(len(old_blocks))]).unsqueeze( - -1 - ) - if update_past_key_values: - merged = torch.cat( - [torch.cat([block, query_tokens[row : row + 1]]) for row, block in enumerate(old_blocks)] - ) - past_key_values.key_cache[0] = merged.reshape(-1, 1, 1) - past_key_values.value_cache[0] = merged.reshape(-1, 1, 1) - return SimpleNamespace(packed_query_sequence=hidden, past_key_values=past_key_values) - - def lm_head(self, hidden): - vocab = 17 - targets = hidden.reshape(-1).long().remainder(vocab) - columns = torch.arange(vocab, dtype=torch.float32, device=hidden.device) - return -(columns.unsqueeze(0) - targets.unsqueeze(1)).abs() - - -class _FakeBagel: - def __init__(self): - self.language_model = _FakeLanguageModel() - self.config = SimpleNamespace(llm_config=SimpleNamespace(num_hidden_layers=1, freeze_und=False)) - - def forward_cache_update_text( - self, - past_key_values, - *, - text_token_lens, - packed_text_ids, - key_values_lens, - **_, - ): - new_blocks = list(packed_text_ids.to(torch.float32).split(text_token_lens.tolist())) - if past_key_values.key_cache[0] is None: - old_blocks = [packed_text_ids.new_zeros(0, dtype=torch.float32) for _ in new_blocks] - else: - old_blocks = list(past_key_values.key_cache[0].reshape(-1).split(key_values_lens.tolist())) - merged = torch.cat([torch.cat([old, new]) for old, new in zip(old_blocks, new_blocks)]) - past_key_values.key_cache[0] = merged.reshape(-1, 1, 1) - past_key_values.value_cache[0] = merged.reshape(-1, 1, 1) - return past_key_values - - -class _FakeBundle: - def __init__(self): - self.model = _FakeBagel() - self.device = torch.device("cpu") - self.new_token_ids = {"bos_token_id": 0, "eos_token_id": 16} - self.transformer = self.model.language_model - - -def _ar_conditions(second_prompt=(7, 8)): - return BagelARConditions( - prompt_splits=[ - [ - {"kind": "text", "ids": torch.tensor([1, 2])}, - {"kind": "text", "ids": torch.tensor([3])}, - ], - [{"kind": "text", "ids": torch.tensor(second_prompt)}], - ] - ) - - -def _ar_stage(forward_batch_size): - bundle = _FakeBundle() - bundle.transformer.eval() - return BagelARStage(model=bundle, forward_batch_size=forward_batch_size) - - -def test_ar_bs1_and_packed_match_tokens_logprobs_and_isolate_samples(): - params = ARSamplingParams( - temperature=0.0, - top_p=1.0, - top_k=0, - max_new_tokens=5, - ) - serial = _ar_stage(1).autoregress(_ar_conditions(), sampling_params=params) - packed = _ar_stage(2).autoregress(_ar_conditions(), sampling_params=params) - - assert torch.equal(serial.tokens, packed.tokens) - assert torch.equal(serial.lengths, packed.lengths) - assert torch.equal(serial.cu_seqlens, packed.cu_seqlens) - torch.testing.assert_close(serial.log_probs, packed.log_probs, rtol=0, atol=0) - - changed_peer = _ar_stage(2).autoregress(_ar_conditions(second_prompt=(12, 13, 14)), sampling_params=params) - first_end = int(packed.cu_seqlens[1]) - assert torch.equal(packed.tokens[:first_end], changed_peer.tokens[:first_end]) - torch.testing.assert_close( - packed.log_probs[:first_end], - changed_peer.log_probs[:first_end], - rtol=0, - atol=0, - ) - - -def test_ar_fbs1_uses_legacy_global_rng_and_packed_uses_generators(monkeypatch): - params = ARSamplingParams( - temperature=0.8, - top_p=0.95, - top_k=8, - max_new_tokens=5, - ) - calls = [] - original_step = BagelARStep.step - - def recording_step(self, logits, *, generators=None): - calls.append(generators) - return original_step(self, logits, generators=generators) - - monkeypatch.setattr(BagelARStep, "step", recording_step) - torch.manual_seed(1234) - _ar_stage(1).autoregress(_ar_conditions(), sampling_params=params) - assert calls and all(generators is None for generators in calls) - - calls.clear() - torch.manual_seed(1234) - _ar_stage(2).autoregress(_ar_conditions(), sampling_params=params) - assert calls and all(generators is not None for generators in calls) - - -def test_ar_vit_input_declines_packing(): - assert BagelARStage._batched_text_ids([[{"kind": "vit", "image": torch.zeros(3, 2, 2)}]]) is None - - -class _FakeStrategy: - def __init__(self): - self.denoise_calls = 0 - - def denoise( - self, - *, - noise_pred, - sample, - prev_sample, - **_, - ): - self.denoise_calls += 1 - mean = sample - 0.25 * noise_pred - output = mean if prev_sample is None else prev_sample - log_prob = -((output - mean) ** 2).mean(dim=(1, 2)) - return output, log_prob, mean - - -def test_diffusion_single_sample_uses_legacy_denoise_path(): - strategy = _FakeStrategy() - step = BagelDiffusionStep() - step.denoise( - strategy, - v_t=torch.ones(3, 4), - x_t=torch.zeros(3, 4), - sigma=torch.tensor(1.0), - sigma_next=torch.tensor(0.5), - sigma_max=torch.tensor(1.0), - eta=0.8, - ) - assert strategy.denoise_calls == 1 - - -def test_diffusion_packed_reduction_matches_bs1_and_isolates_rows(): - step = BagelDiffusionStep() - strategy = _FakeStrategy() - sample = torch.arange(24, dtype=torch.float32).reshape(6, 4) - velocity = torch.arange(24, dtype=torch.float32).flip(0).reshape(6, 4) - kwargs = dict( - sigma=torch.tensor(1.0), - sigma_next=torch.tensor(0.5), - sigma_max=torch.tensor(1.0), - eta=0.8, - ) - - packed = step.denoise( - strategy, - v_t=velocity, - x_t=sample, - n_samples=2, - **kwargs, - ) - serial = [ - step.denoise( - strategy, - v_t=velocity[index * 3 : (index + 1) * 3], - x_t=sample[index * 3 : (index + 1) * 3], - **kwargs, - ) - for index in range(2) - ] - - torch.testing.assert_close(packed[0], torch.cat([value[0] for value in serial])) - torch.testing.assert_close(packed[1], torch.stack([value[1] for value in serial])) - torch.testing.assert_close(packed[2], torch.cat([value[2] for value in serial])) - - changed_velocity = velocity.clone() - changed_velocity[3:] += 1000 - changed = step.denoise( - strategy, - v_t=changed_velocity, - x_t=sample, - n_samples=2, - **kwargs, - ) - torch.testing.assert_close(packed[0][:3], changed[0][:3], rtol=0, atol=0) - torch.testing.assert_close(packed[1][0], changed[1][0], rtol=0, atol=0) - - -def test_diffusion_fixed_per_sequence_rng_matches_serial_steps(monkeypatch): - torch_utils = ModuleType("diffusers.utils.torch_utils") - - def randn_tensor(shape, *, generator, device, dtype): - if isinstance(generator, list): - return torch.cat( - [ - torch.randn((1, *shape[1:]), generator=row_generator, device=device, dtype=dtype) - for row_generator in generator - ] - ) - return torch.randn(shape, generator=generator, device=device, dtype=dtype) - - torch_utils.randn_tensor = randn_tensor - monkeypatch.setitem(sys.modules, "diffusers", ModuleType("diffusers")) - monkeypatch.setitem(sys.modules, "diffusers.utils", ModuleType("diffusers.utils")) - monkeypatch.setitem(sys.modules, "diffusers.utils.torch_utils", torch_utils) - - step = BagelDiffusionStep() - packed_strategy = FlowSDEStrategy() - serial_strategies = [FlowSDEStrategy(), FlowSDEStrategy()] - packed_generators = [torch.Generator().manual_seed(seed) for seed in (11, 29)] - serial_generators = [torch.Generator().manual_seed(seed) for seed in (11, 29)] - packed_sample = torch.arange(24, dtype=torch.float32).reshape(6, 4) / 10 - serial_samples = [packed_sample[:3].clone(), packed_sample[3:].clone()] - schedule = [1.0, 0.8, 0.5, 0.2] - - for current, following in zip(schedule, schedule[1:]): - packed_velocity = packed_sample * 0.1 + 0.25 - packed_sample, packed_logp, _ = step.denoise( - packed_strategy, - v_t=packed_velocity, - x_t=packed_sample, - sigma=torch.tensor(current), - sigma_next=torch.tensor(following), - sigma_max=torch.tensor(0.8), - eta=0.8, - n_samples=2, - generators=packed_generators, - ) - - serial_logps = [] - for index in range(2): - velocity = serial_samples[index] * 0.1 + 0.25 - serial_samples[index], logp, _ = step.denoise( - serial_strategies[index], - v_t=velocity, - x_t=serial_samples[index], - sigma=torch.tensor(current), - sigma_next=torch.tensor(following), - sigma_max=torch.tensor(0.8), - eta=0.8, - generators=[serial_generators[index]], - ) - serial_logps.append(logp) - - torch.testing.assert_close(packed_sample, torch.cat(serial_samples), rtol=0, atol=0) - torch.testing.assert_close(packed_logp, torch.stack(serial_logps), rtol=0, atol=0) - - -def _context(length, start=0): - cache = NaiveCache(1) - if length: - values = torch.arange(start, start + length, dtype=torch.float32).reshape(-1, 1, 1) - cache.key_cache[0] = values - cache.value_cache[0] = values + 100 - return {"kv_lens": [length], "ropes": [length], "past_key_values": cache} - - -def test_diffusion_context_merge_and_pack_fallback_rules(): - merged = BagelDiffusionStage._merge_contexts([_context(2), _context(0), _context(1, 9)]) - assert merged["kv_lens"] == [2, 0, 1] - assert merged["ropes"] == [2, 0, 1] - torch.testing.assert_close( - merged["past_key_values"].key_cache[0].reshape(-1), - torch.tensor([0.0, 1.0, 9.0]), - ) - - same_shape = BagelDiffusionConditions( - gen_contexts=[_context(1), _context(1)], - cfg_text_contexts=[_context(0), _context(0)], - cfg_img_contexts=[_context(1), _context(1)], - prompts=["a", "b"], - image_shapes=[(512, 512), (512, 512)], - ) - mixed_shape = BagelDiffusionConditions( - gen_contexts=same_shape.gen_contexts, - cfg_text_contexts=same_shape.cfg_text_contexts, - cfg_img_contexts=same_shape.cfg_img_contexts, - prompts=same_shape.prompts, - image_shapes=[(512, 512), (384, 640)], - ) - deferred = BagelDiffusionConditions( - prompts=["a", "b"], - image_shapes=[(512, 512), (512, 512)], - ) - - assert BagelDiffusionStage._can_pack_conditions(same_shape) - assert not BagelDiffusionStage._can_pack_conditions(mixed_shape) - assert not BagelDiffusionStage._can_pack_conditions(deferred)