Skip to content
Closed
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
41 changes: 33 additions & 8 deletions backends/vulkan/runtime/graph/ops/glsl/convert.glslh
Original file line number Diff line number Diff line change
Expand Up @@ -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
22 changes: 21 additions & 1 deletion backends/vulkan/runtime/graph/ops/glsl/reduce.glsl
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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.
Expand Down
4 changes: 2 additions & 2 deletions backends/vulkan/runtime/graph/ops/glsl/reduce.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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
17 changes: 12 additions & 5 deletions backends/vulkan/runtime/graph/ops/glsl/reduce_op_defs.glslh
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -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;
}
Expand All @@ -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;
}
Expand All @@ -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;
}
Expand All @@ -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;
}
Expand Down
72 changes: 72 additions & 0 deletions backends/vulkan/test/test_vulkan_dynamic.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Loading