You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
{{ message }}
Repository navigation
[RFC][BAGEL-7B-MoT][CUDA/ROCm] WS1/WS2 kernel roadmap and integration plan #435
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:
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.
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.
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.
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.
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:
The strict objective:
2. Checkpoint fingerprint
BagelForConditionalGeneration,visual_gen+visual_und1e-6; plus per-head Q/K RMSNorm over 1281e6; all latent tokens of one image share one position id*_moe_gentwingenmode: latent rows -> gen weights, text rows (incl. image markers) -> und weights;undmode: all und[markers, latents]queries over[prompt KV, markers, latents]z = 0.3611 * (z_raw - 0.1159)(p, q, c)order (chpwq -> hwpqc)vae2llm64 -> 3584,llm2vae3584 -> 64, with bias[4096, 3584], id =h*64 + wt in [0, 1], no x1000(H/16)*(W/16)latents + 2 markers; 4096 + 2 at 1024², max side 1024t = s*u / (1 + (s-1)*u),u = linspace(1, 0, N), terminal dropped, N-1 Euler stepsx <- x - v*dt(0.4, 1.0],globalrenorm, min 0.0Upstream references:
vllm_omni/diffusion/models/bagel/Upstream runtime status
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.jsonsaysmax_latent_size: 32; the weights hold a 64x64 table. At 512² both give in-range, different position ids.config.jsonsaystimestep_shift: 1.0; the pipeline default is 3.0. At N = 50 that is 30 vs 41 guided steps.globalCFG renorm reduces over the packed batch; per-request combine exists only in step execution.is_causal=Trueis top-left aligned.log_probcomes 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_fingerprintlands these as a versioned profile. Changing any of them is a new profile.scomes from the run config, neverconfig.json.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.logpsums over all latent elements in a fixed order.t = 1first step andt -> 0last step are separate test cases.t), renorm type and min are policy. Renorm is per request.cfg_text_scale <= 1is a separate profile.*_moe_gen,vae2llm,llm2vae,time_embedder), by LoRA or full fine-tune. Frozen und operators still needdXfor marker rows.Required invariances
allclosex_tbyte-equal to rolloutv_t,logp_tTP8 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.
bagel_arch_fingerprintmot_fused_add_rmsnormmot_routed_gemmdXfor und rows,dX+dW/ LoRA for gen rowsmot_qk_norm_ropemot_swiglu_mlpgen_joint_attentionlatent_embed_iovae2llm+ time + position embedding in fixed order,llm2vaecfg_combine_renormflow_sde_step_logpdlogpfull_model_chaintp_invariancesp_invariancecfg_parallel_invariancefsdp_grad_invariancefull_model_chainund_text_logprobvit_navit_connectorvae_encode_conditiont = 0latents in KV)latent_embed_iothink_then_generate_chain5. Recommended claim order
bagel_arch_fingerprintmot_fused_add_rmsnorm,mot_routed_gemm,mot_qk_norm_rope,mot_swiglu_mlpgen_joint_attentionin parallel withlatent_embed_io,cfg_combine_renorm,flow_sde_step_logpfull_model_chaintp_invariance,sp_invariance,cfg_parallel_invariance,fsdp_grad_invariance6. Contribution notes
dtsign and latent layout explicitly.Open questions
globalrenorm, or switch tochannel?If you are interested
Claim one row above and link the implementation PR.