Describe the bug
For models with BFloat16 logits output, onnxruntime-genai returns garbage logits, producing a wrong/incoherent first token and degenerate generation. The identical model exported in Float16 works correctly.
Root cause
src/models/logits.cpp (Logits::Get()) only converts the model's raw logits to float32 when the output type is Float16. For BFloat16 it skips the conversion, then WrapTensor<float>(...) reinterprets the raw 2-byte bf16 bytes as 4-byte float32, corrupting every logit.
Two places are gated on Float16 only:
// ~line 47 — fp32 staging buffer is allocated only for Float16
if (type_ == Ort::TypeToTensorType<Ort::Float16_t>)
logits_of_last_token_fp32_ = OrtValue::CreateTensor<float>(...);
// ~line 73 — Cast to float32 runs only for Float16
if (type_ == Ort::TypeToTensorType<Ort::Float16_t>) {
Cast(*logits_of_last_token, logits_of_last_token_fp32_, *model_.p_device_inputs_, Ort::TypeToTensorType<float>);
logits_of_last_token = logits_of_last_token_fp32_.get();
}
For a bf16 model the Cast never runs, so logits_ ends up wrapping raw bf16 bytes as float.
How I verified
Isolated a bf16 vs fp16 model (Gemma-style decoder, vocab 262144):
- HF bf16 reference, raw ORT
sess.run of the bf16 decoder, and genai's own dumped model-output logits all agree (argmax 496, top-5 [496,886,506,9079,992]).
- genai's
get_logits() for the same bf16 model returns argmax 4539 — a value that does not correspond to any position in the model's actual output. The corruption is purely in the logits read-out, not the model compute or the ORT kernels.
Proposed fix
Treat BFloat16 the same as Float16 in both spots, and add a bf16→f32 conversion path. The existing Cast() helper already falls back to a CPU bf16→f32 cast; I also added a CUDA LaunchBf16ToFp32 kernel so the cast stays on-device.
--- a/src/models/logits.cpp
+++ b/src/models/logits.cpp
- if (type_ == Ort::TypeToTensorType<Ort::Float16_t>)
+ if (type_ == Ort::TypeToTensorType<Ort::Float16_t> || type_ == Ort::TypeToTensorType<Ort::BFloat16_t>)
logits_of_last_token_fp32_ = OrtValue::CreateTensor<float>(...);
...
- // Convert from float16 to float32 if necessary
- if (type_ == Ort::TypeToTensorType<Ort::Float16_t>) {
+ // Convert from float16/bfloat16 to float32 if necessary
+ if (type_ == Ort::TypeToTensorType<Ort::Float16_t> || type_ == Ort::TypeToTensorType<Ort::BFloat16_t>) {
Cast(*logits_of_last_token, logits_of_last_token_fp32_, *model_.p_device_inputs_, Ort::TypeToTensorType<float>);
logits_of_last_token = logits_of_last_token_fp32_.get();
}
Plus a CUDA bf16→f32 case (src/cuda/interface.cpp Cast, src/cuda/model_kernels.cu / kernels.h LaunchBf16ToFp32).
After the fix, the bf16 model's first-token argmax is 496 (matches HF/fp16) and generation is byte-identical to fp16.
Urgency
Medium — blocks deploying any bf16 model via genai. fp16 is an available workaround but loses bf16's wider dynamic range.
Platform / Version
- onnxruntime-genai built from source, CUDA EP, Linux.
- Reproduces independent of EP/kernel dispatch (model compute is correct; only the logits read-out is affected).
I have a tested fix and can open a PR if useful.
Describe the bug
For models with BFloat16 logits output,
onnxruntime-genaireturns garbage logits, producing a wrong/incoherent first token and degenerate generation. The identical model exported in Float16 works correctly.Root cause
src/models/logits.cpp(Logits::Get()) only converts the model's raw logits tofloat32when the output type is Float16. For BFloat16 it skips the conversion, thenWrapTensor<float>(...)reinterprets the raw 2-byte bf16 bytes as 4-byte float32, corrupting every logit.Two places are gated on
Float16only:For a bf16 model the
Castnever runs, sologits_ends up wrapping raw bf16 bytes asfloat.How I verified
Isolated a bf16 vs fp16 model (Gemma-style decoder, vocab 262144):
sess.runof the bf16 decoder, and genai's own dumped model-output logits all agree (argmax 496, top-5[496,886,506,9079,992]).get_logits()for the same bf16 model returns argmax 4539 — a value that does not correspond to any position in the model's actual output. The corruption is purely in the logits read-out, not the model compute or the ORT kernels.Proposed fix
Treat
BFloat16the same asFloat16in both spots, and add a bf16→f32 conversion path. The existingCast()helper already falls back to a CPU bf16→f32 cast; I also added a CUDALaunchBf16ToFp32kernel so the cast stays on-device.Plus a CUDA bf16→f32 case (
src/cuda/interface.cppCast,src/cuda/model_kernels.cu/kernels.hLaunchBf16ToFp32).After the fix, the bf16 model's first-token argmax is 496 (matches HF/fp16) and generation is byte-identical to fp16.
Urgency
Medium — blocks deploying any bf16 model via genai. fp16 is an available workaround but loses bf16's wider dynamic range.
Platform / Version
I have a tested fix and can open a PR if useful.