webgpu : fix handling of infinity values during ARGSORT and TOP_K (#27538)

Co-authored-by: Stanisław Szymczyk <sszymczy@gmail.com>
This commit is contained in:
fairydreaming
2026-08-25 08:08:06 +03:00
committed by GitHub
co-authored by Stanisław Szymczyk
parent f280b26983
commit 5ea87ddad2
@@ -34,11 +34,9 @@ var<uniform> params: Params;
var<workgroup> shmem_idx: array<u32, WG_SIZE>; var<workgroup> shmem_idx: array<u32, WG_SIZE>;
#if ORDER == 0 #if ORDER == 0
#define EXTREME_VALUE 1e30
#define SWAP_COMPARE_UP > #define SWAP_COMPARE_UP >
#define SWAP_COMPARE_DOWN < #define SWAP_COMPARE_DOWN <
#else #else
#define EXTREME_VALUE -1e30
#define SWAP_COMPARE_UP < #define SWAP_COMPARE_UP <
#define SWAP_COMPARE_DOWN > #define SWAP_COMPARE_DOWN >
#endif #endif
@@ -78,11 +76,9 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>,
let dir_up = (lid.x & k) == 0; let dir_up = (lid.x & k) == 0;
let a_idx = shmem_idx[lid.x]; let a_idx = shmem_idx[lid.x];
let b_idx = shmem_idx[ixj]; let b_idx = shmem_idx[ixj];
let a_val = select(EXTREME_VALUE, src[row_base + a_idx], a_idx < params.src_ne0);
let b_val = select(EXTREME_VALUE, src[row_base + b_idx], b_idx < params.src_ne0);
let should_swap = select( let should_swap = select(
(a_val SWAP_COMPARE_DOWN b_val), b_idx >= params.src_ne0 || (a_idx < params.src_ne0 && src[row_base + a_idx] SWAP_COMPARE_DOWN src[row_base + b_idx]),
(a_val SWAP_COMPARE_UP b_val), a_idx >= params.src_ne0 || (b_idx < params.src_ne0 && src[row_base + a_idx] SWAP_COMPARE_UP src[row_base + b_idx]),
dir_up); dir_up);
if (should_swap) { if (should_swap) {
shmem_idx[lid.x] = b_idx; shmem_idx[lid.x] = b_idx;