From 000b4aea0f74d633fcb80407247021c420953b57 Mon Sep 17 00:00:00 2001 From: Mergen Nachin Date: Mon, 28 Sep 2026 18:12:41 -0400 Subject: [PATCH] [Vulkan] Stop clamping per-row reductions and propagate NaN convert.glslh guarded an fp16 clamp with `#if T == float16_t`, but neither side is a macro, so the preprocessor compares 0 with 0 and the clamp was always on. Per-row buffer sum, mean, amax and amin therefore saturated fp32 and int results at 65504. The clamp is removed. Single-dimension texture and per-row buffer amax and amin now propagate NaN, argmax and argmin return the index of the first NaN as ATen does, and mean divides in the accumulator type. FP16 texture output uses explicit nearest-even conversion so rounding and overflow agree with ATen across drivers. Authored with OpenAI Codex; split planned with Claude Code. --- .../runtime/graph/ops/glsl/convert.glslh | 41 ++++++++--- .../vulkan/runtime/graph/ops/glsl/reduce.glsl | 22 +++++- .../vulkan/runtime/graph/ops/glsl/reduce.yaml | 4 +- .../graph/ops/glsl/reduce_op_defs.glslh | 17 +++-- backends/vulkan/test/test_vulkan_dynamic.py | 72 +++++++++++++++++++ 5 files changed, 140 insertions(+), 16 deletions(-) diff --git a/backends/vulkan/runtime/graph/ops/glsl/convert.glslh b/backends/vulkan/runtime/graph/ops/glsl/convert.glslh index b901bc7e9d9..2c3c441b904 100644 --- a/backends/vulkan/runtime/graph/ops/glsl/convert.glslh +++ b/backends/vulkan/runtime/graph/ops/glsl/convert.glslh @@ -13,16 +13,41 @@ #ifdef T -#if T == float16_t - -#define convert_to_T(x) T(clamp(x, -65504, 65504)); - -#else - #define convert_to_T(x) T(x); -#endif // T == float16_t - #endif // T +float round_to_half_rte(float value) { + const uint bits = floatBitsToUint(value); + const uint exponent = (bits >> 23u) & 0xffu; + if (exponent == 0xffu) { + return value; + } + + uint result = 0u; + if (exponent >= 143u) { + result = 0x7c00u; + } else if (exponent >= 102u) { + const bool normal = exponent >= 113u; + const uint significand = (bits & 0x7fffffu) | (normal ? 0u : 0x800000u); + const uint shift = normal ? 13u : 126u - exponent; + result = (normal ? (exponent - 112u) << 10u : 0u) + (significand >> shift); + const uint remainder = significand & ((1u << shift) - 1u); + const uint halfway = 1u << (shift - 1u); + result += uint(remainder > halfway || (remainder == halfway && (result & 1u) != 0u)); + } + + // packHalf2x16 does not guarantee round-to-nearest-even on every driver. + result |= (bits >> 16u) & 0x8000u; + return unpackHalf2x16(result).x; +} + +vec4 round_to_half_rte(vec4 value) { + return vec4( + round_to_half_rte(value.x), + round_to_half_rte(value.y), + round_to_half_rte(value.z), + round_to_half_rte(value.w)); +} + #endif // CONVERT_GLSLH diff --git a/backends/vulkan/runtime/graph/ops/glsl/reduce.glsl b/backends/vulkan/runtime/graph/ops/glsl/reduce.glsl index 029e3b16756..f37a0fc773a 100644 --- a/backends/vulkan/runtime/graph/ops/glsl/reduce.glsl +++ b/backends/vulkan/runtime/graph/ops/glsl/reduce.glsl @@ -49,6 +49,7 @@ shared vec4 shared_vecs[MAX_NTHREADS]; #include "indexing_utils.h" #include "indexing.glslh" +#include "convert.glslh" int tid_to_smi(const ivec2 tid) { return tid.x + tid.y * NWORKERS; @@ -84,7 +85,26 @@ int tid_to_smi(const ivec2 tid) { #define UPDATE_ACCUM(accum, new_val) ${UPDATE_ACCUM} // Useful for operators such as mean which want to perform a final calculation // with the accumulator. -#define POSTPROCESS(accum) ${POSTPROCESS} +$if DTYPE == "half": + #define POSTPROCESS(accum) round_to_half_rte(${POSTPROCESS}) +$else: + #define POSTPROCESS(accum) ${POSTPROCESS} + +float max_propagate_nan(float a, float b) { + return isnan(a) ? a : (isnan(b) ? b : max(a, b)); +} + +vec4 max_propagate_nan(vec4 a, vec4 b) { + return mix(mix(max(a, b), b, isnan(b)), a, isnan(a)); +} + +float min_propagate_nan(float a, float b) { + return isnan(a) ? a : (isnan(b) ? b : min(a, b)); +} + +vec4 min_propagate_nan(vec4 a, vec4 b) { + return mix(mix(min(a, b), b, isnan(b)), a, isnan(a)); +} /* * Computes reduction where the reduction dim is orthogonal to the packed dim. diff --git a/backends/vulkan/runtime/graph/ops/glsl/reduce.yaml b/backends/vulkan/runtime/graph/ops/glsl/reduce.yaml index 21a7132b8db..f85d8405b2d 100644 --- a/backends/vulkan/runtime/graph/ops/glsl/reduce.yaml +++ b/backends/vulkan/runtime/graph/ops/glsl/reduce.yaml @@ -21,9 +21,9 @@ reduce: POSTPROCESS: (accum / tin_sizes[reduce_dim]) - NAME: amax INIT_ACCUM: first_val - UPDATE_ACCUM: max(accum, new_val) + UPDATE_ACCUM: max_propagate_nan(accum, new_val) POSTPROCESS: accum - NAME: amin INIT_ACCUM: first_val - UPDATE_ACCUM: min(accum, new_val) + UPDATE_ACCUM: min_propagate_nan(accum, new_val) POSTPROCESS: accum diff --git a/backends/vulkan/runtime/graph/ops/glsl/reduce_op_defs.glslh b/backends/vulkan/runtime/graph/ops/glsl/reduce_op_defs.glslh index e5f61da7586..338ad955760 100644 --- a/backends/vulkan/runtime/graph/ops/glsl/reduce_op_defs.glslh +++ b/backends/vulkan/runtime/graph/ops/glsl/reduce_op_defs.glslh @@ -15,6 +15,11 @@ struct Accum { uint count; }; +bool is_earlier_nan(ACCUM_T val, uint idx, const Accum accum) { + return isnan(float(val)) && + (!isnan(float(accum.val)) || idx < accum.idx); +} + void init_accum(out Accum accum, T val, uint idx) { accum.val = ACCUM_T(val); accum.idx = idx; @@ -40,14 +45,14 @@ void merge_accum_sum(inout Accum accum, const Accum other) { } void postprocess_accum_mean(inout Accum accum) { - accum.val /= T(accum.count); + accum.val /= ACCUM_T(accum.count); } // Amax (maximum value) void update_accum_amax(inout Accum accum, T in_val, uint idx) { ACCUM_T val = ACCUM_T(in_val); - if (val > accum.val) { + if (val > accum.val || is_earlier_nan(val, idx, accum)) { accum.val = val; accum.idx = idx; } @@ -58,7 +63,7 @@ void update_accum_amax(inout Accum accum, T in_val, uint idx) { } void merge_accum_amax(inout Accum accum, const Accum other) { - if (other.val > accum.val) { + if (other.val > accum.val || is_earlier_nan(other.val, other.idx, accum)) { accum.val = other.val; accum.idx = other.idx; } @@ -72,7 +77,7 @@ void merge_accum_amax(inout Accum accum, const Accum other) { void update_accum_amin(inout Accum accum, T in_val, uint idx) { ACCUM_T val = ACCUM_T(in_val); - if (val < accum.val) { + if (val < accum.val || is_earlier_nan(val, idx, accum)) { accum.val = val; accum.idx = idx; } @@ -83,7 +88,9 @@ void update_accum_amin(inout Accum accum, T in_val, uint idx) { } void merge_accum_amin(inout Accum accum, const Accum other) { - if (other.count > 0 && (accum.count == 0 || other.val < accum.val)) { + if (other.count > 0 && + (accum.count == 0 || other.val < accum.val || + is_earlier_nan(other.val, other.idx, accum))) { accum.val = other.val; accum.idx = other.idx; } diff --git a/backends/vulkan/test/test_vulkan_dynamic.py b/backends/vulkan/test/test_vulkan_dynamic.py index b671337d2a0..089c14eae4b 100644 --- a/backends/vulkan/test/test_vulkan_dynamic.py +++ b/backends/vulkan/test/test_vulkan_dynamic.py @@ -161,6 +161,78 @@ def test_gelu_with_singleton_dimensions(self): edge = self._lower(model, (x,), storage=storage) self._run(edge, model, [(x,)], atol=5e-6, rtol=5e-6) + def test_buffer_reduction_range(self): + class Reduce(torch.nn.Module): + def __init__(self, op): + super().__init__() + self.op = op + + def forward(self, x): + return self.op(x, dim=-1, keepdim=True) + + for op in (torch.sum, torch.mean, torch.amax): + with self.subTest(op=op): + width, value = (20000, 4) if op == torch.sum else (8, 80000) + x = torch.tensor([value, -value], dtype=torch.float32)[:, None].repeat( + 1, width + ) + model = Reduce(op) + edge = self._lower(model, (x,), storage=VkStorageType.BUFFER) + self._run(edge, model, [(x,)], atol=0, rtol=0) + + def test_reduction_special_values(self): + class Reduce(torch.nn.Module): + def __init__(self, op): + super().__init__() + self.op = op + + def forward(self, x): + return self.op(x, dim=-1, keepdim=True) + + for dtype in (torch.float32, torch.float16): + x = torch.tensor( + [ + [40000] * 9, + [-40000] * 9, + [1, torch.nan, 2, 3, torch.nan, 4, 5, 6, 7], + ], + dtype=dtype, + ) + for op in (torch.sum, torch.mean, torch.amax, torch.amin): + for storage in (VkStorageType.TEXTURE_3D, VkStorageType.BUFFER): + with self.subTest(dtype=dtype, op=op, storage=storage): + model = Reduce(op) + edge = self._lower(model, (x,), storage=storage) + self._run(edge, model, [(x,)], atol=0, rtol=0, equal_nan=True) + + x = torch.tensor([4, -4], dtype=torch.float16)[:, None].repeat(1, 70000) + model = Reduce(torch.mean) + edge = self._lower(model, (x,), storage=VkStorageType.BUFFER) + self._run(edge, model, [(x,)], atol=0, rtol=0) + + def test_argreduce_first_nan(self): + class Reduce(torch.nn.Module): + def __init__(self, op): + super().__init__() + self.op = op + + def forward(self, x): + return self.op(x, dim=-1, keepdim=True) + + for dtype in (torch.float32, torch.float16): + x = torch.tensor( + [ + [1, torch.nan, 2, torch.nan, 3, 4, 5], + [1, 2, 3, 4, 5, torch.nan, torch.nan], + ], + dtype=dtype, + ) + for op in (torch.argmax, torch.argmin): + with self.subTest(dtype=dtype, op=op): + model = Reduce(op) + edge = self._lower(model, (x,), storage=VkStorageType.BUFFER) + self._run(edge, model, [(x,)], atol=0, rtol=0) + if __name__ == "__main__": unittest.main()