From 507c122c3f29188498b6452a9f8334d23101577c Mon Sep 17 00:00:00 2001 From: Mergen Nachin Date: Mon, 28 Sep 2026 18:26:58 -0400 Subject: [PATCH] [Vulkan] Match ATen special values in pow.Tensor_Scalar power_of only special-cased zero bases and used a 1e-5 tolerance to detect odd integer exponents, so NaN, infinities, signed zeros and negative bases with fractional exponents gave wrong results. It now follows ATen case by case, including the FP32 sqrt and rsqrt path for exponents of plus or minus 0.5. The requested tensor dtype survives device storage emulation, so FP16 keeps its distinct ATen special-value semantics on SwiftShader. Explicit nearest-even output conversion preserves the FP16 conformance cases on hardware GPUs. Authored with OpenAI Codex; split planned with Claude Code. --- .../vulkan/runtime/api/containers/Tensor.cpp | 4 ++ .../vulkan/runtime/api/containers/Tensor.h | 6 ++ backends/vulkan/runtime/graph/ComputeGraph.h | 6 ++ .../graph/ops/glsl/binary_op_defs.glslh | 70 +++++++++++++------ .../graph/ops/glsl/binary_scalar_buffer.glsl | 10 ++- .../graph/ops/glsl/binary_scalar_texture.glsl | 7 ++ .../runtime/graph/ops/impl/BinaryScalarOp.cpp | 2 +- backends/vulkan/test/test_vulkan_dynamic.py | 23 ++++++ 8 files changed, 104 insertions(+), 24 deletions(-) diff --git a/backends/vulkan/runtime/api/containers/Tensor.cpp b/backends/vulkan/runtime/api/containers/Tensor.cpp index 748f65baa9d..85cfdbed6e7 100644 --- a/backends/vulkan/runtime/api/containers/Tensor.cpp +++ b/backends/vulkan/runtime/api/containers/Tensor.cpp @@ -886,6 +886,7 @@ vTensor::vTensor( const vkapi::VulkanImage* external_image, const vkapi::VulkanBuffer* external_buffer) : dtype_(get_effective_scalar_type(context, dtype, memory_layout)), + original_dtype_(dtype), packed_dim_info_(calculate_packed_dim_info(memory_layout, storage_type)), // Calculate tensor metadata sizes_(sizes.begin(), sizes.end()), @@ -962,6 +963,7 @@ vTensor::vTensor( const utils::GPUMemoryLayout memory_layout, const utils::AxisMapLayout axis_map_layout) : dtype_(vkapi::element_scalartype(image.format())), + original_dtype_(dtype_), packed_dim_info_( calculate_packed_dim_info(memory_layout, utils::kTexture3D)), // Calculate tensor metadata @@ -996,6 +998,7 @@ vTensor::vTensor( vTensor::vTensor(vTensor& other) : dtype_(other.dtype_), + original_dtype_(other.original_dtype_), packed_dim_info_{other.packed_dim_info_}, // Copy tensor size metadata sizes_(other.sizes_.begin(), other.sizes_.end()), @@ -1021,6 +1024,7 @@ vTensor::vTensor( const std::vector& sizes, const std::vector& dim_order) : dtype_(other.dtype_), + original_dtype_(other.original_dtype_), packed_dim_info_(other.packed_dim_info_), // Copy tensor size metadata sizes_(sizes.begin(), sizes.end()), diff --git a/backends/vulkan/runtime/api/containers/Tensor.h b/backends/vulkan/runtime/api/containers/Tensor.h index 8e6b7a4a133..07a0ddc890a 100644 --- a/backends/vulkan/runtime/api/containers/Tensor.h +++ b/backends/vulkan/runtime/api/containers/Tensor.h @@ -354,6 +354,8 @@ class vTensor final { // Whether the tensor has elements of type float, int, etc. vkapi::ScalarType dtype_; + // Requested dtype before device-specific storage emulation. + vkapi::ScalarType original_dtype_; // Information about packed dimension padding and block packing PackedDimInfo packed_dim_info_; // sizes of the tensor in NCHW dimension order @@ -519,6 +521,10 @@ class vTensor final { return dtype_; } + inline vkapi::ScalarType original_dtype() const { + return original_dtype_; + } + /* * Provide a "best guess" of a memory layout that can be used to construct a * tensor with similar layout metadata (i.e. strides, axis_map, etc.) as this diff --git a/backends/vulkan/runtime/graph/ComputeGraph.h b/backends/vulkan/runtime/graph/ComputeGraph.h index eb01e3abf5e..0cb22089d49 100644 --- a/backends/vulkan/runtime/graph/ComputeGraph.h +++ b/backends/vulkan/runtime/graph/ComputeGraph.h @@ -371,6 +371,12 @@ class ComputeGraph final { vkapi::ScalarType dtype_of(const ValueRef idx) const; + inline vkapi::ScalarType original_dtype_of(const ValueRef idx) const { + const Value& value = values_.at(idx); + return value.isTensor() ? value.toConstTensor().original_dtype() + : dtype_of(idx); + } + vkapi::ScalarType get_staging_dtype_for(const ValueRef idx) const; inline const utils::ivec3& logical_limits_of(const ValueRef idx) const { diff --git a/backends/vulkan/runtime/graph/ops/glsl/binary_op_defs.glslh b/backends/vulkan/runtime/graph/ops/glsl/binary_op_defs.glslh index 83ba3e4e0d3..a755503f173 100644 --- a/backends/vulkan/runtime/graph/ops/glsl/binary_op_defs.glslh +++ b/backends/vulkan/runtime/graph/ops/glsl/binary_op_defs.glslh @@ -9,16 +9,6 @@ #ifndef BINARY_OP_DEFS_GLSLH #define BINARY_OP_DEFS_GLSLH -// -// Power operation that handles negative and zero bases -// -// In GLSL, pow(x, y) is undefined for x < 0. This function provides -// a safe implementation that: -// - Handles x == 0 (returns 0 for y > 0, returns 1 for y == 0) -// - Handles x < 0 by using absolute value and preserving sign for odd integer exponents -// - Uses standard pow() for x > 0 -// - // Operands are evaluated in the promoted compute type COMPUTE_T (and its vector // form COMPUTE_VEC4_T), which the including shader sets from // get_higher_precision_dtype(DTYPE, SCALAR_VALUE_TYPE) so mixed tensor/scalar @@ -27,23 +17,61 @@ // Scalar overload COMPUTE_T power_of(COMPUTE_T x, COMPUTE_T y) { - if (x == 0.0) { - // Handle 0^y: 0^0 = 1, 0^y = 0 for y > 0 - return (y == 0.0) ? COMPUTE_T(1.0) : COMPUTE_T(0.0); + const float base = float(x); + float exponent = float(y); + const float infinity = uintBitsToFloat(0x7f800000u); + const float nan = uintBitsToFloat(0x7fc00000u); + + if (input_is_half) { + exponent = round_to_half_rte(exponent); + } + + if (exponent == 0.0 || base == 1.0) { + return COMPUTE_T(1.0); + } + if (isnan(base) || isnan(exponent)) { + return COMPUTE_T(nan); } - // Use absolute value to avoid undefined behavior - float result = pow(abs(float(x)), float(y)); + if (!input_is_half) { + // ATen uses sqrt/rsqrt for FP32 scalar exponents, but generic pow for half. + if (abs(exponent) == 0.5) { + if (base < 0.0) { + return COMPUTE_T(nan); + } + if (base == 0.0) { + const bool negative_zero = (floatBitsToUint(base) & 0x80000000u) != 0u; + return exponent > 0.0 ? x : COMPUTE_T(negative_zero ? -infinity : infinity); + } + return COMPUTE_T(exponent > 0.0 ? sqrt(base) : 1.0 / sqrt(base)); + } + } - // For negative bases with odd integer exponents, preserve the negative sign - if (x < 0.0) { - float int_y = round(float(y)); - if (abs(float(y) - int_y) < 1e-5 && int(int_y) % 2 == 1) { - result = -result; + const float magnitude = abs(base); + if (isinf(exponent)) { + if (magnitude == 1.0) { + return COMPUTE_T(1.0); } + return COMPUTE_T((magnitude > 1.0) == (exponent > 0.0) ? infinity : 0.0); + } + + const bool integral_exponent = trunc(exponent) == exponent; + const bool odd_exponent = integral_exponent && mod(abs(exponent), 2.0) == 1.0; + const bool negative = (floatBitsToUint(base) & 0x80000000u) != 0u && odd_exponent; + + float result; + if (base == 0.0) { + result = exponent > 0.0 ? 0.0 : infinity; + } else if (isinf(base)) { + result = exponent > 0.0 ? infinity : 0.0; + } else if (base < 0.0 && !integral_exponent) { + result = nan; + } else { + // GLSL pow requires a positive base. + result = pow(magnitude, exponent); } - return COMPUTE_T(result); + return COMPUTE_T(negative ? -result : result); } #ifdef COMPUTE_VEC4_T diff --git a/backends/vulkan/runtime/graph/ops/glsl/binary_scalar_buffer.glsl b/backends/vulkan/runtime/graph/ops/glsl/binary_scalar_buffer.glsl index 9def3666870..168dae1ce4f 100644 --- a/backends/vulkan/runtime/graph/ops/glsl/binary_scalar_buffer.glsl +++ b/backends/vulkan/runtime/graph/ops/glsl/binary_scalar_buffer.glsl @@ -40,6 +40,7 @@ ${define_active_storage_type(STORAGE)} layout(std430) buffer; #include "indexing.glslh" +#include "convert.glslh" $if IS_COMPARISON_OP: ${layout_declare_tensor(B, "w", "t_out", "uint8", STORAGE)} @@ -56,6 +57,7 @@ layout(push_constant) uniform restrict Block { }; layout(local_size_x_id = 0, local_size_y_id = 1, local_size_z_id = 2) in; +layout(constant_id = 3) const bool input_is_half = false; #include "dispatch.glslh" @@ -68,6 +70,10 @@ void main() { return; } - t_out[out_bufi] = - OUT_T(op(COMPUTE_T(t_in[out_bufi]), COMPUTE_T(scalar_value))); + COMPUTE_T value = COMPUTE_T(op(COMPUTE_T(t_in[out_bufi]), COMPUTE_T(scalar_value))); + $if not IS_COMPARISON_OP and DTYPE in ("float", "half"): + if (input_is_half) { + value = COMPUTE_T(round_to_half_rte(float(value))); + } + t_out[out_bufi] = OUT_T(value); } diff --git a/backends/vulkan/runtime/graph/ops/glsl/binary_scalar_texture.glsl b/backends/vulkan/runtime/graph/ops/glsl/binary_scalar_texture.glsl index acd523a67ba..440e81d095b 100644 --- a/backends/vulkan/runtime/graph/ops/glsl/binary_scalar_texture.glsl +++ b/backends/vulkan/runtime/graph/ops/glsl/binary_scalar_texture.glsl @@ -42,6 +42,7 @@ ${define_active_storage_type(STORAGE)} layout(std430) buffer; #include "indexing.glslh" +#include "convert.glslh" $if IS_COMPARISON_OP: ${layout_declare_tensor(B, "w", "t_out", "uint8", STORAGE)} @@ -58,6 +59,7 @@ layout(push_constant) uniform restrict Block { }; layout(local_size_x_id = 0, local_size_y_id = 1, local_size_z_id = 2) in; +layout(constant_id = 3) const bool input_is_half = false; $if not IS_COMPARISON_OP: #include "binary_op_defs.glslh" @@ -73,5 +75,10 @@ void main() { VEC4_OUT_T out_texel = VEC4_OUT_T( op(COMPUTE_VEC4_T(in_texel), COMPUTE_VEC4_T(scalar_value))); + $if not IS_COMPARISON_OP and DTYPE in ("float", "half"): + if (input_is_half) { + out_texel = round_to_half_rte(out_texel); + } + imageStore(t_out, pos, out_texel); } diff --git a/backends/vulkan/runtime/graph/ops/impl/BinaryScalarOp.cpp b/backends/vulkan/runtime/graph/ops/impl/BinaryScalarOp.cpp index 57bfff019ab..352c2987180 100644 --- a/backends/vulkan/runtime/graph/ops/impl/BinaryScalarOp.cpp +++ b/backends/vulkan/runtime/graph/ops/impl/BinaryScalarOp.cpp @@ -112,7 +112,7 @@ void add_binary_scalar_op_node( // Push Constants push_constants, // Specialization Constants - {}, + {int32_t(graph.original_dtype_of(out) == vkapi::kHalf)}, // Resize Args {}, // Resizing Logic diff --git a/backends/vulkan/test/test_vulkan_dynamic.py b/backends/vulkan/test/test_vulkan_dynamic.py index 8e10060b426..3b049925c06 100644 --- a/backends/vulkan/test/test_vulkan_dynamic.py +++ b/backends/vulkan/test/test_vulkan_dynamic.py @@ -715,6 +715,29 @@ def forward(self, x): self.assertEqual(_vulkan_graphs(edge), []) self._run(edge, model, [(x,)]) + def test_power_special_values(self): + class Power(torch.nn.Module): + def __init__(self, exponent): + super().__init__() + self.exponent = exponent + + def forward(self, x): + return torch.pow(x, self.exponent) + + for dtype in (torch.float32, torch.float16): + x = torch.tensor( + [-torch.inf, -10000, -4, -0.0, 0.0, 1, 10000, torch.inf, torch.nan], + dtype=dtype, + ).repeat(3, 1) + for exponent in (-3, -0.5, 0, 0.5, 2, 3, torch.inf, -torch.inf): + for storage in (VkStorageType.TEXTURE_3D, VkStorageType.BUFFER): + with self.subTest(dtype=dtype, exponent=exponent, storage=storage): + model = Power(exponent) + edge = self._lower(model, (x,), storage=storage) + self._run( + edge, model, [(x,)], equal_nan=True, check_signed_zero=True + ) + def test_fp16_scalar_rounding(self): class CreateTensor(torch.nn.Module): def __init__(self, value, scalar):