Skip to content

[RFC][BAGEL-7B-MoT][CUDA/ROCm] WS1/WS2 kernel roadmap and integration plan #435

Description

@zhangj1an

Status: Proposed

Target: post-v0.1.0 community roadmap

Related: #386, #434

Upstream checkpoint pinned for this draft: ByteDance-Seed/BAGEL-7B-MoT@5019f57.


1. Motivation

BAGEL is a unified understanding + generation model. One Qwen2.5-7B-shaped decoder holds two weight sets (Mixture-of-Transformers, MoT) that share self-attention: text and ViT tokens use the understanding expert, VAE-latent tokens use the generation expert. Images are generated by rectified flow inside the same decoder, conditioned on a causally prefilled prompt KV cache.

This adds three things RL-Kernel does not cover today:

  1. Token-routed dense weights. Every GEMM and norm picks its weight set per token by a fixed index set. No learned router, but a row's bytes must not depend on the modality mix of the pack.
  2. A diffusion policy inside an LLM decoder. Latent tokens query the prompt KV non-causally at every step. The policy logprob is the ODE→SDE transition density, as in [RFC][Qwen-Image][CUDA/ROCm] WS1/WS2 kernel roadmap, ablation matrix and integration plan #386, but with BAGEL's own schedule, sign and timestep convention.
  3. Three-branch renormalised CFG. Up to three branches, two scales, a timestep window and a norm-ratio rescale. The reference renorm reduces over the whole packed batch.

The strict objective:

Given the same weights, prompt tokens, initial latents, schedule, CFG parameters and SDE noise, training and rollout must execute the same declared arithmetic contract and produce exactly equal per-step velocities v_t and log-probabilities logp_t in strict mode.

  • Phase A - text-to-image: generation-path operators, full-trajectory parity, TP / SP / CFG-parallel invariance.
  • Phase B - text-output policy, image-conditioned editing, and think-then-generate. Opens after Phase A closeout.

2. Checkpoint fingerprint

Item Value
Architecture BagelForConditionalGeneration, visual_gen + visual_und
Parameters 14.61B BF16 (LLM 14.14B, of which gen expert 6.53B; ViT 0.40B) + FP32 VAE 0.084B; ~7B active per token
Decoder 28 layers, hidden 3584, GQA 28 / 4 heads, D = 128
Projections Q/K/V 3584 -> 3584/512/512 with bias; O 3584 -> 3584 no bias
Norms RMSNorm eps 1e-6; plus per-head Q/K RMSNorm over 128
MLP SwiGLU 3584 -> 2x18944 -> 3584, no bias
RoPE 1-D NeoX, theta 1e6; all latent tokens of one image share one position id
MoT weights Q/K/V/O, Q/K norms, layer norms, MLP and final norm each have a *_moe_gen twin
Routing gen mode: latent rows -> gen weights, text rows (incl. image markers) -> und weights; und mode: all und
Attention causal prompt prefill; non-causal [markers, latents] queries over [prompt KV, markers, latents]
VAE FLUX-style, 16 ch, 8x downsample, z = 0.3611 * (z_raw - 0.1159)
Latent layout 2x2 patch, 64 values/token in (p, q, c) order (chpwq -> hwpqc)
Latent heads vae2llm 64 -> 3584, llm2vae 3584 -> 64, with bias
Latent positions frozen 2-D sin-cos table [4096, 3584], id = h*64 + w
Timestep embedding 256-ch cos-then-sin, SiLU MLP 256 -> 3584 -> 3584; input t in [0, 1], no x1000
Tokens (H/16)*(W/16) latents + 2 markers; 4096 + 2 at 1024², max side 1024
Schedule t = s*u / (1 + (s-1)*u), u = linspace(1, 0, N), terminal dropped, N-1 Euler steps x <- x - v*dt
CFG defaults text 4.0, image 1.5, interval (0.4, 1.0], global renorm, min 0.0
Phase B only SigLIP ViT (1152, patch 14, 26 of 27 layers used), connector 1152 -> 3584 -> 3584 tanh-GELU

Upstream references:

Upstream runtime status

  • vllm-omni serves BAGEL (single-stage and Thinker + DiT) with CFG-parallel, Ulysses / Ring SP, step execution and opt-in trajectory output. It is the initial rollout runtime.
  • Upstream BAGEL trains with FSDP and supports freeze_und. No VIME / Megatron provider exists; treat trainer binding as new integration work.

Reference hazards

These are plausible implementation errors that raise nothing:

  • config.json says max_latent_size: 32; the weights hold a 64x64 table. At 512² both give in-range, different position ids.
  • config.json says timestep_shift: 1.0; the pipeline default is 3.0. At N = 50 that is 30 vs 41 guided steps.
  • BAGEL runs N-1 steps; the Lance subclass runs N.
  • global CFG renorm reduces over the packed batch; per-request combine exists only in step execution.
  • Sequential, CFG-parallel, SP and step-batched paths pack rows differently.
  • QK norm runs in FP32, then RoPE and Q/K in BF16.
  • Initial noise is drawn on CPU in one path and regenerated on device in another.
  • Cached prefill needs a bottom-right causal mask; SDPA is_causal=True is top-left aligned.
  • Rollout log_prob comes from a pluggable scheduler; there is no in-tree BAGEL SDE.

3. Numerical contract

The WS1 numerical standard remains authoritative: fixed accumulator precision and reduction order, no Split-K / Stream-K / split-KV / atomics without a contracted merge tree, no TF32 or fast math, casts only at declared boundaries, fail closed on unsupported geometry.

BAGEL-specific rules

bagel_arch_fingerprint lands these as a versioned profile. Changing any of them is a new profile.

  • Schedule built in FP64, stored and traced in FP32. s comes from the run config, never config.json.
  • SDE step and logp in FP32 with sigma = t, dt = t_i - t_{i+1} > 0. The sign is the opposite of [RFC][Qwen-Image][CUDA/ROCm] WS1/WS2 kernel roadmap, ablation matrix and integration plan #386. logp sums over all latent elements in a fixed order.
  • t = 1 first step and t -> 0 last step are separate test cases.
  • CFG branch set, scales, interval (compared on FP32 t), renorm type and min are policy. Renorm is per request. cfg_text_scale <= 1 is a separate profile.
  • Initial and transition noise use an explicit per-sample generator with fixed device, dtype and consumption order.
  • Phase A trains only the gen expert (*_moe_gen, vae2llm, llm2vae, time_embedder), by LoRA or full fine-tune. Frozen und operators still need dX for marker rows.

Required invariances

Level Required comparison
Operator accuracy candidate vs independent FP64 reference, allclose
Batch / packing same bytes under request count, pack position and padding
Modality mix text / latent rows unchanged as the other modality's row count varies
CFG branch same bytes batched, sequential or on separate ranks
Step recompute trainer step from stored x_t byte-equal to rollout v_t, logp_t
Prompt context KV byte-equal across full, chunked and variable-length batched prefill
Distributed TP2 / TP4, SP2 / SP4, CFG2 / CFG3, FSDP gradients byte-equal to single rank

TP8 is unsupported (28 heads). CUDA and ROCm each require exact parity within a pinned profile; cross-platform byte equality is reported, not assumed.


4. Work-item table

Status: OPEN -> IN PROGRESS -> IN REVIEW -> MERGED.

To claim a task, put your handle in the GitHub column and open a PR. A row is complete only when implementation, independent reference, invariance tests and benchmarks against the native path land together.

Work item What it does Track Platform Reuse / dependency GitHub PR Status
bagel_arch_fingerprint Revision, MoT weight map, routing, geometry, latent / VAE metadata, per-layer trace, versioned profile foundation CUDA + ROCm New @zhangj1an IN PROGRESS
mot_fused_add_rmsnorm Residual add + routed RMSNorm, incl. final norm WS1 CUDA + ROCm Extend #201 OPEN
mot_routed_gemm Routed Q/K/V + bias and O; dX for und rows, dX + dW / LoRA for gen rows WS1 CUDA + ROCm Extend #180/#396 OPEN
mot_qk_norm_rope Routed FP32 per-head QK RMSNorm + NeoX RoPE, declared RoPE precision WS1 CUDA + ROCm Extend #350/#228 OPEN
mot_swiglu_mlp Routed gate/up + SiLU-mul + down, single cast WS1 CUDA + ROCm Reuse #280/#180 OPEN
gen_joint_attention Non-causal GQA 28/4/128 over prompt KV + generated rows; varlen branch masks; bottom-right cached prefill WS1 CUDA + ROCm Extend #240/#246/#319 OPEN
latent_embed_io Pack / unpack, vae2llm + time + position embedding in fixed order, llm2vae WS1 CUDA + ROCm Extend #410; different layout OPEN
cfg_combine_renorm Three-branch combine, interval gate, per-request renorm, FP32 fixed norm tree WS1 CUDA + ROCm New OPEN
flow_sde_step_logp Shifted schedule, FP32 SDE transition and logp, dlogp WS1 CUDA + ROCm Extend #386 OPEN
full_model_chain 28-layer real-weight trajectory, per-step parity, first-drift by layer / operator WS1 closeout CUDA + ROCm Extend #315 @zhangj1an IN PROGRESS
tp_invariance TP2 / TP4 for both experts, fixed all-reduce WS2 CUDA + ROCm Extend #310 OPEN
sp_invariance Ulysses / Ring SP2 / SP4, replicated markers + KV, fixed gather; one TP x SP case WS2 CUDA + ROCm Extend #235 OPEN
cfg_parallel_invariance CFG2 / CFG3 with fixed broadcast, branch placement and gather; one SP x CFG case WS2 CUDA + ROCm New OPEN
fsdp_grad_invariance FSDP / HSDP gradient parity for gen weights or LoRA WS2 CUDA + ROCm New; full_model_chain OPEN
und_text_logprob Und prefill / decode + LM head selected-token logp Phase B CUDA + ROCm Reuse #204/#243/#336 OPEN
vit_navit_connector 26-layer SigLIP, block-diagonal packing, connector Phase B CUDA + ROCm New OPEN
vae_encode_condition VAE encode + pack for image context (t = 0 latents in KV) Phase B CUDA + ROCm Extend latent_embed_io OPEN
think_then_generate_chain Joint AR-text + image-transition logp over one KV cache Phase B CUDA + ROCm All Phase B rows OPEN

5. Recommended claim order

  1. bagel_arch_fingerprint
  2. mot_fused_add_rmsnorm, mot_routed_gemm, mot_qk_norm_rope, mot_swiglu_mlp
  3. gen_joint_attention in parallel with latent_embed_io, cfg_combine_renorm, flow_sde_step_logp
  4. full_model_chain
  5. tp_invariance, sp_invariance, cfg_parallel_invariance, fsdp_grad_invariance
  6. Phase B rows

6. Contribution notes

  • Do not claim a routed row with one expert. Both weight sets and the modality-mix test are required.
  • Do not reuse [RFC][Qwen-Image][CUDA/ROCm] WS1/WS2 kernel roadmap, ablation matrix and integration plan #386 kernels where the contract differs. Test timestep scaling, dt sign and latent layout explicitly.
  • Do not use decoded images as the first comparator. Compare velocities, logp and latents first.
  • Strict mode fails closed on batch-coupled renorm, a position grid that disagrees with the table, unpinned RNG, and unsupported TP / SP layouts.

Open questions

  1. Gen-expert LoRA or full fine-tune?
  2. Transport hash-checked rollout KV to the trainer, or rebuild it?
  3. Keep per-request global renorm, or switch to channel?
  4. Keep the reference BF16 RoPE, or define an FP32 profile?

If you are interested

Claim one row above and link the implementation PR.

Activity

  1. changed the title [-][RFC][BAGEL-7B-MoT][CUDA/ROCm] Strict Training–Rollout Parity[/-] [+][RFC][BAGEL-7B-MoT][CUDA/ROCm] WS1/WS2 kernel roadmap and integration plan[/+] on Sep 19, 2026
  2. added
    platform: cudaSpecific optimizations or bugs in NVIDIA graphics cards (such as FlashInfer, TMA optimizations)
    platform: rocmSpecific tasks specific to AMD graphics cards (such as CK, bpreshuffle/FA)
    multimodalFeatures, bugs, or optimizations specific to multimodal support.
    on Sep 19, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

Labels

multimodalFeatures, bugs, or optimizations specific to multimodal support.new-modelnext-phaseplatform: cudaSpecific optimizations or bugs in NVIDIA graphics cards (such as FlashInfer, TMA optimizations)platform: rocmSpecific tasks specific to AMD graphics cards (such as CK, bpreshuffle/FA)

Type

No type

Projects

No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions