diff --git a/common/sampling.cpp b/common/sampling.cpp index 256ac161e20f..e935f29d93a8 100644 --- a/common/sampling.cpp +++ b/common/sampling.cpp @@ -574,6 +574,23 @@ llama_token common_sampler_sample(struct common_sampler * gsmpl, struct llama_co gsmpl->set_logits(ctx, idx); + // engine#123: HRX can hand back all-NaN router logits when the driver migrates a page + // behind in-flight work (the engine#140-class nondeterminism). NaN is never a legitimate + // logit - unlike -inf, which masking legitimately uses - so refuse to sample from it + // instead of silently decoding a wrong-but-plausible token. + { + const float * logits = llama_get_logits_ith(ctx, idx); + if (logits != nullptr) { + const int32_t n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(llama_get_model(ctx))); + for (int32_t i = 0; i < n_vocab; ++i) { + if (std::isnan(logits[i])) { + LOG_ERR("%s: HRX returned NaN logits at vocab index %d - refusing to sample a silently wrong token (engine#123)\n", __func__, (int) i); + GGML_ABORT("HRX: NaN logits (engine#123)"); + } + } + } + } + // Check if a backend sampler has already sampled a token in which case we // return that token id directly. { diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-flash-attention.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-flash-attention.cpp index 461d55474117..14cb02ded57a 100644 --- a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-flash-attention.cpp +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-flash-attention.cpp @@ -6,6 +6,7 @@ #include #include +#include #include #include #include @@ -13,6 +14,25 @@ namespace ggml::hrx { namespace { +// Test knob for the decode-split partial transients' alignment. Production is 4096; +// the engine#123/#140 rig oracle used 256, a layout that exposes the divergence. +// Defaults to 4096 so production behaviour is unchanged unless explicitly requested. +static size_t decode_split_partial_alignment() { + static const size_t value = []() -> size_t { + const char * env = std::getenv("GGML_HRX_FA_PARTIAL_ALIGN"); + if (env == nullptr || *env == '\0') { + return 4096u; + } + char * end = nullptr; + const long parsed = std::strtol(env, &end, 10); + if (end == env || parsed <= 0 || parsed > (1 << 20)) { + return 4096u; + } + return static_cast(parsed); + }(); + return value; +} + static constexpr KernelCatalogRef kFlashAttentionF32F16WmmaKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_flash_attention_f32_f16_wmma"); static constexpr KernelCatalogRef kFlashAttentionDecodeSplitNextQ8Kernel = @@ -578,11 +598,11 @@ static bool match_flash_attention_decode_split_next_q8_dispatch(const DispatchMa const ValueId q8_output = match_value(context, dispatch_match, 4); dispatch_match.transients.push_back( - { partial_max, "common.decode.flash_attention.partial_max", partial_scalar_bytes, 4096 }); + { partial_max, "common.decode.flash_attention.partial_max", partial_scalar_bytes, decode_split_partial_alignment() }); dispatch_match.transients.push_back( - { partial_sum, "common.decode.flash_attention.partial_sum", partial_scalar_bytes, 4096 }); + { partial_sum, "common.decode.flash_attention.partial_sum", partial_scalar_bytes, decode_split_partial_alignment() }); dispatch_match.transients.push_back( - { partial_output, "common.decode.flash_attention.partial_output", partial_output_bytes, 4096 }); + { partial_output, "common.decode.flash_attention.partial_output", partial_output_bytes, decode_split_partial_alignment() }); dispatch_match.transients.push_back( { q8_output, "common.decode.flash_attention.next_q8_output", q8_output_bytes, 4096 }); dispatch_match.completion_counter_requests.push_back({ diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/router_top8_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/router_top8_f32.loom index 7c2f77fae093..ae4113298146 100644 --- a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/router_top8_f32.loom +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/router_top8_f32.loom @@ -112,12 +112,20 @@ template.def<@qwen3_moe.router.top8.row> device requires [#target.subgroup.size< %unnormalized_weight = scalar.expf %selected_delta : f32 %selected_sum = kernel.subgroup.reduce %unnormalized_weight : f32 %route_weight = scalar.divf %unnormalized_weight, %selected_sum : f32 + // engine#123: when no lane of the row found an ordered candidate, the row's router logits + // are entirely NaN. The argmax seed is a finite -FLT_MAX, so the softmax below still + // produces a plausible uniform 1/route_count, which routes the token to a wrong-but-valid + // expert and hides the corruption completely. Publish the row's own NaN instead, so the + // corruption reaches the host logits and is caught loudly rather than silently. + %row_had_ordered = scalar.cmpf ogt, %selected_max, %negative_large : f32 + %raw_candidate = vector.extract %initial_logits[%c0] : vector<[%experts_per_lane]xf32> -> f32 + %published_weight = scf.select %row_had_ordered, %route_weight, %raw_candidate : f32 %lane_publishes_weight = index.cmp ult, %lane, %route_count : index %publishes_weight = scalar.andi %valid_token, %lane_publishes_weight : i1 scf.if %publishes_weight { %route_weight_token_base = index.mul %safe_token, %route_count : index %route_weight_index = index.add %route_weight_token_base, %lane : index - view.store %route_weight, %route_weights_view[%route_weight_index] : f32, view<[%route_weight_storage_count]xf32> + view.store %published_weight, %route_weights_view[%route_weight_index] : f32, view<[%route_weight_storage_count]xf32> } template.return }