From c16424f5620d635b847ac782d39939e80a97d959 Mon Sep 17 00:00:00 2001 From: agent Date: Sun, 27 Sep 2026 09:36:18 -0300 Subject: [PATCH] hrx: make the all-NaN router-logit case loud instead of a silent wrong expert engine#123's residual, after #25 stopped the 0x7FFFFFFF sentinel from faulting: when the driver migrates a page behind in-flight HRX work (the engine#140-class nondeterminism), a row's router logits go all-NaN. The argmax seed is a finite -FLT_MAX, so the row's softmax still produced a plausible uniform 1/route_count and the token decoded with a valid-but-wrong expert, silently. - router_top8_f32.loom: when no lane of the row found an ordered candidate, publish the row's own NaN instead of that masked uniform weight, so the corruption reaches the logits rather than being hidden behind a plausible one. - common/sampling.cpp: refuse to sample NaN logits - log and abort loudly. NaN is never a legitimate logit, unlike -inf, which masking legitimately uses. - dispatch-flash-attention.cpp: GGML_HRX_FA_PARTIAL_ALIGN selects the decode-split partial transients' alignment (default 4096, production unchanged); 256 reproduces the engine#123/#140 rig oracle without editing stress literals. Verified on the MoE repro (Qwen3-Coder-30B-A3B-Instruct Q4_K_M, 2113 ctx, fresh server per sample, under memory pressure): before, 6/13 samples silently decoded ' Paris???????????????'; after, the corrupted samples abort with "HRX returned NaN logits ... refusing to sample a silently wrong token (engine#123)". --- common/sampling.cpp | 17 ++++++++++++ .../common/dispatch-flash-attention.cpp | 26 ++++++++++++++++--- .../qwen_moe/qwen3_moe/router_top8_f32.loom | 10 ++++++- 3 files changed, 49 insertions(+), 4 deletions(-) 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 }