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
14 changes: 14 additions & 0 deletions backends/vulkan/op_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -812,6 +812,20 @@ def register_reduce_cpp_ops():
)


@update_features(exir_ops.edge.aten.any.dim)
def register_any_dim():
return OpFeatures(
inputs_storage=utils.ANY_TEXTURE,
inputs_dtypes=utils.BOOL_T,
supports_resize=True,
supports_highdim=True,
are_node_inputs_supported_fn=lambda node: (
utils.ndim_of(node.args[0]) > 0 and is_reduce_node_supported(node)
),
pick_io_storage_fn=pick_storage_for_reduce,
)


# =============================================================================
# ArgReduce.cpp
# =============================================================================
Expand Down
15 changes: 8 additions & 7 deletions backends/vulkan/runtime/graph/ops/glsl/reduce.glsl
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@

#define PRECISION ${PRECISION}
#define VEC4_T ${texel_load_type(DTYPE, STORAGE)}
#define T ${texel_load_component_type(DTYPE, STORAGE)}

${define_active_storage_type(STORAGE)}

Expand Down Expand Up @@ -45,7 +46,7 @@ layout(constant_id = 6) const int NWORKERS = 4;
#define MAX_NTHREADS 256


shared vec4 shared_vecs[MAX_NTHREADS];
shared VEC4_T shared_vecs[MAX_NTHREADS];

#include "indexing_utils.h"
#include "indexing.glslh"
Expand Down Expand Up @@ -123,7 +124,7 @@ void reduce_nonpacked_dim(
// behaviour that hangs some GPUs. They still take a shared memory slot, but
// it is one that no in-bounds group aggregates over, so what they leave in it
// is never read.
vec4 accum = vec4(0);
VEC4_T accum = VEC4_T(0);
if (in_bounds) {
scan_pos[reduce_dim] = 0;
accum = INIT_ACCUM(load_texel(tin, scan_pos));
Expand Down Expand Up @@ -195,10 +196,10 @@ void reduce_packed_dim(
// behaviour that hangs some GPUs. They still take a shared memory slot, but
// it is one that no in-bounds group aggregates over, so what they leave in it
// is never read.
vec4 accum = vec4(0);
VEC4_T accum = VEC4_T(0);
if (in_bounds) {
scan_pos[reduce_dim] = 0;
accum = INIT_ACCUM(vec4(load_texel(tin, scan_pos).x));
accum = INIT_ACCUM(VEC4_T(load_texel(tin, scan_pos).x));

// Partially accumulate over elements i, i + NWORKERS, i + 2*NWORKERS, ...
// of the reduction row
Expand All @@ -212,7 +213,7 @@ void reduce_packed_dim(
// padding elements are ignored
if (scan_pos[reduce_dim] == safe_idx(tin_limits, reduce_dim) - 1 &&
nspill > 0) {
const vec4 intex = load_texel(tin, scan_pos);
const VEC4_T intex = load_texel(tin, scan_pos);
for (int i = 0; i < nspill; i++) {
accum.x = UPDATE_ACCUM(accum.x, intex[i]);
}
Expand All @@ -233,13 +234,13 @@ void reduce_packed_dim(
}
// Each element of the texel is itself a partial maximum; iterate over the
// texel to find the actual maximum
float accum_final = accum.x;
T accum_final = accum.x;
[[unroll]] for (int i = 1; i < 4; i++) {
accum_final = UPDATE_ACCUM(accum[i], accum_final);
}

scan_pos[reduce_dim] = tid.x;
write_texel(tout, scan_pos, POSTPROCESS(vec4(accum_final, 0, 0, 0)));
write_texel(tout, scan_pos, POSTPROCESS(VEC4_T(accum_final, 0, 0, 0)));
}
}

Expand Down
3 changes: 3 additions & 0 deletions backends/vulkan/runtime/graph/ops/glsl/reduce.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,9 @@ reduce:
- VALUE: float
shader_variants:
- NAME: sum
- NAME: any_uint8
DTYPE: uint8
UPDATE_ACCUM: max(accum, new_val)
- NAME: mean
POSTPROCESS: (accum / tin_sizes[reduce_dim])
- NAME: amax
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,10 @@ reduce_per_row_buffer:
- VALUE: int32
shader_variants:
- NAME: sum_per_row_buffer
- NAME: any_per_row_buffer_uint8
DTYPE: uint8
UPDATE_ACCUM_FN: update_accum_amax
MERGE_ACCUM_FN: merge_accum_amax
- NAME: mean_per_row_buffer
POSTPROCESS_ACCUM_FN: postprocess_accum_mean
- NAME: amax_per_row_buffer
Expand Down
12 changes: 12 additions & 0 deletions backends/vulkan/runtime/graph/ops/impl/Reduce.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -417,11 +417,23 @@ DEFINE_REDUCE_FN(mean, 4)
DEFINE_REDUCE_FN(amax, 3)
DEFINE_REDUCE_FN(amin, 3)

void any_dim(ComputeGraph& graph, const std::vector<ValueRef>& args) {
if (graph.is_buffer_storage(args[0])) {
VK_CHECK_COND(
normalize(
graph.extract_scalar<int64_t>(args[1]), graph.dim_of(args[0])) ==
graph.dim_of(args[0]) - 1);
return add_reduce_per_row_node(graph, args[0], args[2], args[3], "any");
}
return add_reduce_node(graph, args[0], args[1], args[3], "any");
}

REGISTER_OPERATORS {
VK_REGISTER_OP(aten.sum.dim_IntList, sum);
VK_REGISTER_OP(aten.mean.dim, mean);
VK_REGISTER_OP(aten.amax.default, amax);
VK_REGISTER_OP(aten.amin.default, amin);
VK_REGISTER_OP(aten.any.dim, any_dim);
}

} // namespace vkcompute
64 changes: 62 additions & 2 deletions backends/vulkan/test/test_vulkan_dynamic.py
Original file line number Diff line number Diff line change
Expand Up @@ -144,6 +144,27 @@ def _run(
)
)

def test_partition_any_unsupported_inputs(self):
class AnyDim(torch.nn.Module):
def forward(self, x):
return torch.any(x, dim=0, keepdim=True)

for x in (
torch.tensor(True),
torch.tensor([0, 1], dtype=torch.int32),
torch.tensor([0, 1], dtype=torch.uint8),
torch.tensor([0.0, 1.0]),
):
with self.subTest(shape=x.shape, dtype=x.dtype):
edge = to_edge_transform_and_lower(
export(AnyDim(), (x,)),
partitioner=[VulkanPartitioner({"require_dynamic_shapes": True})],
)
self.assertNotIn(
torch.ops.higher_order.executorch_call_delegate,
[node.target for node in edge.exported_program().graph.nodes],
)

def test_dynamic_gelu(self):
for approximate in ("none", "tanh"):
for storage in (VkStorageType.TEXTURE_3D, VkStorageType.BUFFER):
Expand Down Expand Up @@ -249,6 +270,41 @@ def forward(self, x):
)
self._run(edge, model, inputs, atol=0, rtol=0)

def test_dynamic_any_dim(self):
class AnyDim(torch.nn.Module):
def __init__(self, dim, keepdim):
super().__init__()
self.dim = dim
self.keepdim = keepdim

def forward(self, x):
return torch.any(x, dim=self.dim, keepdim=self.keepdim)

for dim, keepdim in ((-1, True), (-1, False), (1, True)):
for storage in (VkStorageType.TEXTURE_3D, VkStorageType.BUFFER):
with self.subTest(dim=dim, keepdim=keepdim, storage=storage):
model = AnyDim(dim, keepdim)
inputs = []
for s in (16, 3, 31, 2, 16):
x = torch.zeros(2, s, 5, dtype=torch.bool)
if s != 3:
x[0, s // 2, 1] = True
x[1, -1, 3] = True
x[1, 0, 4] = True
inputs.append((x,))
seq = Dim("s", min=2, max=32)
unsupported = dim == 1 and storage == VkStorageType.BUFFER
edge = self._lower(
model,
inputs[0],
({1: seq},),
storage,
fully_delegated=not unsupported,
)
if unsupported:
self.assertEqual(_vulkan_graphs(edge), [])
self._run(edge, model, inputs, atol=0, rtol=0)

def test_dynamic_logical_not(self):
class LogicalNot(torch.nn.Module):
def forward(self, x):
Expand Down Expand Up @@ -701,7 +757,7 @@ def __init__(self, op, dim):
def forward(self, x):
return self.op(x, dim=self.dim, keepdim=True)

for op in (torch.sum, torch.mean, torch.amax):
for op in (torch.any, torch.sum, torch.mean, torch.amax):
for batch, dim, supported in (
(1, 0, False),
(2, 0, False),
Expand All @@ -712,7 +768,11 @@ def forward(self, x):
):
with self.subTest(op=op, batch=batch, dim=dim):
values = torch.arange(batch * 3 * 4 * 5).reshape(batch, 3, 4, 5)
x = -((values * 37 + 11) % values.numel() + 1).float() / 7
x = (
values % 7 == 0
if op == torch.any
else -((values * 37 + 11) % values.numel() + 1).float() / 7
)
model = Reduce(op, dim)
edge = self._lower(model, (x,), fully_delegated=supported)
if not supported:
Expand Down
Loading