Skip to content

fix(tts): cast bf16 tables explicitly where the CUDA overlay demotes - #2093

Merged
inureyes merged 2 commits into
mainfrom
fix/issue-2087-cuda-dtype-promotion
Oct 2, 2026
Merged

inureyes merged 2 commits into
mainfrom
fix/issue-2087-cuda-dtype-promotion

Conversation

@inureyes

@inureyes inureyes commented Oct 2, 2026

Copy link
Copy Markdown
Member

Summary

Five of the six CUDA lib-test failures share one cause: mlxcel's own CUDA overlay of MLX's promotion table (src/lib/mlx-cpp/patches-cuda/dtype.cpp) resolves bf16 with f32 to bf16 on every device of a CUDA build. It is not an MLX bug at the pin and not an mlxcel-core wrapper. The overlay is deliberate (#636, single-dtype bf16 decode graph) and stays, so the code that needed f32 now casts explicitly. The sixth failure (gelu) was already fixed on main by #2080 and passes on CUDA unchanged.

  • Production fix (TTS): RvqCodebooks::depthsum_embedding and MogHead::infer gathered rows from stored bf16 tables and combined them with f32, so on CUDA the depth-sum and the MoG mu came back bf16. Both now cast the gathered rows (exact, bit-identical off CUDA). Reverting only these casts makes depthsum_ignores_mask_index and mog_infer_shapes_are_finite_with_guidance fail on CUDA again.
  • Probe test mixed_bf16_f32_promotion_follows_the_build pins add, matmul and fast_rms_norm dtype for f32 against bf16: f32 without cuda, bf16 with it. The issue asked for f32 on both builds; that cannot hold while the overlay is kept, so the probe pins the contract that does hold.
  • Tests: the f32_weights and perception references spell upstream promotion out as an in-graph astype on CUDA; the MoG test loads weights through promoted_subset as RvqEarTtsModel does.
  • Technical report EN+KO added.

Test plan

  • GB10 --release --features cuda, under gpu-lock, --test-threads=1: audio::f32_weights 4, audio::fastconformer 14, models::nemotron_voicechat::tts 13, models::gemma3_backbone 5, all pass
  • CPU-only Linux build (no GPU feature, separate target dir): same filters, same counts, all pass. Not MLXCEL_DEVICE=cpu, since the overlay is compiled in.
  • cargo clippy --release --features cuda --lib --tests -- -D warnings and cargo fmt --check clean

Not validated here:

  • The plain Linux CPU lib-test build does not link on main: turbo/kv_inplace_write.cpp (perf(cache): write the decode KV row in place instead of copying the cache #1961) references mlx::core::copy_gpu_inplace, undefined without a GPU backend. The CPU run used a linker wrapper ignoring that one unresolved symbol.
  • The EAR-TTS backbone still meets f32 activations with stored bf16 tensors (Gemma norm weights, bos_emb, null_emb, audio_prompt_projection_w), so by the probed rule it runs bf16 on CUDA where the reference runs f32. Not measured end to end; no VoiceChat checkpoint on this host.

Closes #2087

inureyes added a commit that referenced this pull request Oct 2, 2026
@inureyes inureyes added status:done Completed type:bug Bug fixes, error corrections, or issue resolutions priority:medium Medium priority labels Oct 2, 2026
mlxcel's CUDA overlay of MLX's promotion table (src/lib/mlx-cpp/patches-cuda/dtype.cpp, kept for the single-dtype bf16 decode graph of #636) resolves bf16 with f32 to bf16 on every device of a CUDA build. Five lib tests and two Nemotron VoiceChat TTS sites assumed upstream's f32: the RVQ depth-sum and the MoG head's gathered proj_mus/low_mat slabs returned bf16 on CUDA. Both now cast the gathered rows explicitly (exact, bit-identical off CUDA).

A probe test pins the per-build rule for add, matmul and fast_rms_norm (f32 off CUDA, bf16 with it). The f32_weights and perception tests build their reference with an explicit in-graph cast on CUDA, and the MoG test loads weights through promoted_subset as the model does. The gelu failure was already fixed by #2080.

Validated on GB10 (--features cuda) and a CPU-only build: f32_weights, fastconformer, nemotron_voicechat::tts and gemma3_backbone all pass; clippy clean.

Closes #2087
@inureyes
inureyes force-pushed the fix/issue-2087-cuda-dtype-promotion branch from 7dec006 to b77d006 Compare October 2, 2026 08:07
@inureyes
inureyes merged commit 324a83d into main Oct 2, 2026
48 of 49 checks passed
@inureyes
inureyes deleted the fix/issue-2087-cuda-dtype-promotion branch October 2, 2026 08:23
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

priority:medium Medium priority status:done Completed type:bug Bug fixes, error corrections, or issue resolutions

Projects

None yet

Development

Successfully merging this pull request may close these issues.

fix: six lib tests fail under --features cuda on main

1 participant