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..e4a4147d601 100644 --- a/backends/vulkan/test/test_vulkan_dynamic.py +++ b/backends/vulkan/test/test_vulkan_dynamic.py @@ -715,6 +715,40 @@ 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, + 2.0001, + 3, + 2049, + 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):