Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions backends/vulkan/runtime/api/containers/Tensor.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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()),
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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()),
Expand All @@ -1021,6 +1024,7 @@ vTensor::vTensor(
const std::vector<int64_t>& sizes,
const std::vector<int64_t>& 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()),
Expand Down
6 changes: 6 additions & 0 deletions backends/vulkan/runtime/api/containers/Tensor.h
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
6 changes: 6 additions & 0 deletions backends/vulkan/runtime/graph/ComputeGraph.h
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down
70 changes: 49 additions & 21 deletions backends/vulkan/runtime/graph/ops/glsl/binary_op_defs.glslh
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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)}
Expand All @@ -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"

Expand All @@ -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);
}
Original file line number Diff line number Diff line change
Expand Up @@ -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)}
Expand All @@ -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"
Expand All @@ -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);
}
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
34 changes: 34 additions & 0 deletions backends/vulkan/test/test_vulkan_dynamic.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
Loading