fix(tts): cast bf16 tables explicitly where the CUDA overlay demotes - #2093
Merged
Merged
Conversation
inureyes
added a commit
that referenced
this pull request
Oct 2, 2026
3 tasks done
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
force-pushed
the
fix/issue-2087-cuda-dtype-promotion
branch
from
October 2, 2026 08:07
7dec006 to
b77d006
Compare
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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 anmlxcel-corewrapper. 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.RvqCodebooks::depthsum_embeddingandMogHead::infergathered rows from stored bf16 tables and combined them with f32, so on CUDA the depth-sum and the MoGmucame back bf16. Both now cast the gathered rows (exact, bit-identical off CUDA). Reverting only these casts makesdepthsum_ignores_mask_indexandmog_infer_shapes_are_finite_with_guidancefail on CUDA again.mixed_bf16_f32_promotion_follows_the_buildpinsadd,matmulandfast_rms_normdtype for f32 against bf16: f32 withoutcuda, 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.f32_weightsand perception references spell upstream promotion out as an in-graphastypeon CUDA; the MoG test loads weights throughpromoted_subsetasRvqEarTtsModeldoes.Test plan
--release --features cuda, undergpu-lock,--test-threads=1:audio::f32_weights4,audio::fastconformer14,models::nemotron_voicechat::tts13,models::gemma3_backbone5, all passMLXCEL_DEVICE=cpu, since the overlay is compiled in.cargo clippy --release --features cuda --lib --tests -- -D warningsandcargo fmt --checkcleanNot validated here:
turbo/kv_inplace_write.cpp(perf(cache): write the decode KV row in place instead of copying the cache #1961) referencesmlx::core::copy_gpu_inplace, undefined without a GPU backend. The CPU run used a linker wrapper ignoring that one unresolved symbol.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