Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 17 additions & 0 deletions common/sampling.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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.
{
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -6,13 +6,33 @@

#include <cmath>
#include <cstdint>
#include <cstdlib>
#include <sstream>
#include <string>
#include <utility>

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<size_t>(parsed);
}();
return value;
}

static constexpr KernelCatalogRef kFlashAttentionF32F16WmmaKernel =
GGML_HRX_KERNEL_REF("loom_libs", "ggml_flash_attention_f32_f16_wmma");
static constexpr KernelCatalogRef kFlashAttentionDecodeSplitNextQ8Kernel =
Expand Down Expand Up @@ -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({
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -112,12 +112,20 @@ template.def<@qwen3_moe.router.top8.row> device requires [#target.subgroup.size<
%unnormalized_weight = scalar.expf<afn> %selected_delta : f32
%selected_sum = kernel.subgroup.reduce<addf> %unnormalized_weight : f32
%route_weight = scalar.divf<reassoc|nnan|ninf|nsz|afn> %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
}
Expand Down
Loading