Adaptive Reward Routing post-trains joint audio-video diffusion models by recomputing, each round, both which tokens/layers reward gradients flow through and how competing rewards are weighted, keeping user-defined priorities as a floor while adding conflict-aware residual corrections.
Suppose you have a model that generates a video and its matching audio from one text prompt (e.g., a barking dog where the bark lines up with the mouth moving). You want it to be good at four different things at once: video quality, audio quality, text-to-media alignment, and audio-video sync. These objectives disagree. Pushing one reward can quietly wreck another.
The standard fix is reward-guided RL on the diffusion model, using multiple reward signals. Prior work on exactly this setting, OmniNFT, already recognized that in a two-stream (separate audio and video branches, connected by Cross-attention) model, you need to route each reward’s gradient to the branch it actually supervises, and localize it to specific layers. But OmniNFT decides that routing once, based on the pretrained model, and then freezes it. The authors show (via ablating cross-modal layers in the base vs. fine-tuned checkpoints) that the layers carrying cross-modal information shift during fine-tuning. A frozen routing map goes stale. Separately, methods like Marble adapt reward weights based on gradient geometry, but can drive a weak-but-important reward’s weight to zero because its gradients look noisy early on. The authors show this literally happening to the audio-video sync reward.
So the gap: nobody adapts both where updates land and how rewards are combined, while still respecting what the user said they cared about.
The method sits on top of DiffusionNFT, a forward-process RL formulation for diffusion models that avoids a learned value function by comparing each sample to its peers for the same prompt. There are two moving parts.
Part 1: routing updates to the right tokens and layers. The honest way to measure “how much does audio-to-video cross-attention at layer L matter right now?” is to turn it off and see how much the output changes. That’s too expensive to do every step. Instead, the authors use a free quantity already computed in the forward pass: the norm of the pre-gate cross-attention output at each layer and token. High norm = this path is carrying a lot of information right now. They aggregate this two ways: average over layers to get a per-token weight (upweight the loss on cross-modally influential tokens), and average over tokens to get a per-layer coefficient. The per-layer coefficient is used as a soft gradient detachment knob: forward pass is unchanged, but backward gradients through weakly-coupled layers get scaled down. Validation: the proxy’s layer ranking matches true ablation ranking at Spearman correlation ~0.97, and a proxy frozen at initialization decays to ~0.4 correlation during training.
Part 2: coordinating rewards without crushing weak ones. Each reward is probed only through the branch it actually supervises (video rewards through the video branch, cross-modal rewards through both). Within each branch, a Marble-style solver produces a conflict-aware coefficient. Then, after a warm-up period, the final weight is a convex mix: (1-κ)·prior + κ·(rescaled gradient coefficient). The user-specified prior sets a nonzero floor so a weak objective can’t be deleted; the gradient term adapts to current conflicts.
for round in training:
samples = rollout(policy)
advantages = {k: normalize(reward_k(samples)) for k in rewards}
token_w = norm99(mean_over_layers(cross_attn_response))
layer_a = (1 - norm(mean_over_tokens(response))) ** (1/tau)
if round >= warmup and refresh_due:
gamma = marble_per_branch(reward_grads_through_own_branch)
w = prior if round < warmup else (1-k)*prior + k*rescale(gamma)
A_branch = sum(w[m,k] * advantages[k] for k in branch_rewards[m])
loss = token_weighted_nft_loss(A_branch, token_w)
backprop_with_layer_gated_kv(loss, layer_a)
Evaluated on JavisBench (10,140 prompts), on two LTX-2 backbones (19B and 22B), against four baselines: no post-training, Group Reward-Decoupled Policy Optimization (GDPO) (fixed weights), Marble (gradient-only weights), and OmniNFT (static modality routing).
•
The full method wins 9 of 10 metrics on each backbone. The headline synchronization metric DeSync (lower is better) drops from the base model’s 0.604 to 0.341 on LTX-2, versus 0.390 for OmniNFT and 0.671 for Group Reward-Decoupled Policy Optimization (GDPO) (which actually makes sync worse while improving video quality).
•
Training-dynamics plot: baselines make lopsided progress (one reward goes up, another flatlines or decays). The full method lifts all five component rewards together. The authors frame this as evidence that gains come from coordinated optimization, not from sacrificing one objective.
•
Ablations isolate the four pieces. Token weighting and layer scaling each help individually and combine additively on the routing side. On the weighting side, branch-aware conflict estimation alone underperforms, but adding the residual-to-prior mix and the warm-up each add gains. The paper notes intermediate configs don’t improve every metric monotonically; the full stack is where the balance lands.
•
Proxy validation (Sec. 5.4): disabling the top-10% proxy-scored tokens’ cross-modal inputs changes the output 1.62× (A2V) and 1.74× (V2A) more than disabling a random 10%. So the free proxy really is pointing at functionally important tokens, not just high-magnitude noise.
•
If you’re doing multi-reward RL post-training on a two-stream multimodal diffusion model, the single most portable idea here is the residual-to-prior reward weighting: keep the user’s weights as a floor, let gradient-conflict methods only perturb them. This directly addresses the failure mode where gradient-only weighters (like Marble) silently zero out an objective the user said mattered. Worth trying even without the routing machinery.
•
The pre-gate cross-attention norm as a free influence proxy is a cheap diagnostic even outside RL. If you want to know which layers in a cross-attention-coupled model are doing cross-modal work right now, you can read it off the forward pass without ablation sweeps. The authors validate this against actual intervention, which is the right check.
•
The soft gradient detachment via α·sg(x) + (1-α)·x trick is a clean way to scale gradients through a path without changing the forward value. Useful anywhere you want differentiable control over backward flow through specific components.
•
Caveat on scope: the routing piece needs identifiable per-modality tokens and accessible cross-modal interaction responses. The authors say reward coordination is architecture-independent but routing needs the two-stream structure (or similar). If you’re on a fully-unified single-stream model with no separable cross-modal pathway, only the weighting half transfers cleanly.
•
The paper does not mention a code release in the supplied text.
•
Everything is evaluated on one benchmark family (JavisBench) and two closely related backbones from the same model line. Generalization to other joint audio-video architectures is asserted but not shown here, except for a brief mention of a unified single-stream backbone in the conclusion.
•
The reward-coordination comparison inherits the limitations of the reward models themselves (VideoAlign, HPSv3, AudioBox Aesthetics, CLAP, DeSync). The authors acknowledge no unified audio-video reward exists; optimizing against this specific bundle doesn’t guarantee human-perceived quality improvements.
•
Standard deviations across 3 seeds are reported and are small relative to the deltas vs. baselines, but 3 seeds is still 3 seeds. Treat specific decimal gaps between adjacent configs with appropriate skepticism.
•
The method adds non-trivial machinery (per-round proxy aggregation, periodic Marble solves, warm-up schedule, smoothing). The paper does not quantify the training-time overhead vs. baselines in the supplied text.
•
The qualitative claim that it preserves identity across frames better is from four cherry-picked examples. Take it as illustrative, not as a measured identity-consistency result.