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()