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
19 changes: 16 additions & 3 deletions backends/vulkan/op_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -720,9 +720,8 @@ def is_reduce_node_supported_by_general_impl(node: torch.fx.Node) -> bool:
if isinstance(dims_reduced, (list, tuple)) and not 1 <= len(dims_reduced) <= 2:
return False

keepdim = get_keepdim_setting(node)
# keepdim = False is not supported yet for general implementation
if isinstance(keepdim, bool) and not keepdim:
# any.dim can repack the reduced texture after removing the reduction axis.
if not get_keepdim_setting(node) and node.target != exir_ops.edge.aten.any.dim:
return False

if utils.ndim_of(node.args[0]) == 4:
Expand Down Expand Up @@ -812,6 +811,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
20 changes: 12 additions & 8 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 All @@ -251,7 +252,10 @@ void main() {
gl_LocalInvocationID[reduce_dim],
gl_LocalInvocationID[group_dim]);

const bool in_bounds = all(lessThan(scan_pos, tin_limits));
// Reducing an empty dimension still produces one output element.
ivec3 out_limits = tin_limits;
out_limits[reduce_dim] = 1;
const bool in_bounds = all(lessThan(scan_pos, out_limits));

// reduce_dim and packed_dim are specialization constants, so this branch is
// uniform across the work group and safe to take around a barrier.
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
40 changes: 40 additions & 0 deletions backends/vulkan/runtime/graph/ops/impl/Reduce.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@

#include <executorch/backends/vulkan/runtime/graph/ops/impl/Common.h>
#include <executorch/backends/vulkan/runtime/graph/ops/impl/Staging.h>
#include <executorch/backends/vulkan/runtime/graph/ops/impl/View.h>

#include <executorch/backends/vulkan/runtime/graph/ops/impl/utils/TensorUtils.h>
#include <executorch/backends/vulkan/runtime/graph/ops/utils/ShaderNameUtils.h>
Expand Down Expand Up @@ -419,11 +420,50 @@ DEFINE_REDUCE_FN(mean, 4)
DEFINE_REDUCE_FN(amax, 3)
DEFINE_REDUCE_FN(amin, 3)

void resize_squeezed_reduce_node(
ComputeGraph* graph,
const std::vector<ArgGroup>& args,
const std::vector<ValueRef>& resize_args) {
const ValueRef out = args.at(0).refs.at(0);
const ValueRef in = args.at(1).refs.at(0);
std::vector<int64_t> sizes = graph->sizes_of(in);
const int64_t dim = normalize(
graph->extract_scalar<int64_t>(resize_args.at(0)), sizes.size());
sizes.erase(sizes.begin() + dim);
graph->virtual_resize(out, sizes);
}

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");
}
if (graph.extract_scalar<bool>(args[2])) {
return add_reduce_node(graph, args[0], args[1], args[3], "any");
}

std::vector<int64_t> sizes = graph.sizes_of(args[0]);
sizes.at(normalize(graph.extract_scalar<int64_t>(args[1]), sizes.size())) = 1;
TmpTensor reduced(
&graph,
sizes,
graph.dtype_of(args[0]),
graph.storage_type_of(args[0]),
graph.estimate_memory_layout_of(args[0]));
add_reduce_node(graph, args[0], args[1], reduced, "any");
return add_view_copy_node(
graph, reduced, args[3], {args[1]}, resize_squeezed_reduce_node);
}

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
150 changes: 148 additions & 2 deletions backends/vulkan/test/test_vulkan_dynamic.py
Original file line number Diff line number Diff line change
Expand Up @@ -149,6 +149,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 @@ -254,6 +275,127 @@ def forward(self, x):
)
self._run(edge, model, inputs, atol=0, rtol=0)

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

def forward(self, x):
mask = x > 0
if self.keepdim is None:
reduced = torch.any(mask, dim=-1)
else:
reduced = torch.any(mask, dim=-1, keepdim=self.keepdim)
return torch.logical_not(reduced), torch.any(reduced, dim=-1)

inputs = [
((torch.arange(2 * s * 5).reshape(2, s, 5) % 17).float() - 14,)
for s in (16, 3, 31, 2, 16)
]
for keepdim in (None, False):
with self.subTest(keepdim=keepdim):
model = AnyDim(keepdim)
edge = self._lower(
model,
inputs[0],
({1: Dim("s", min=2, max=32)},),
)
for graph in _vulkan_graphs(edge):
for value in graph.values:
tensor = value.value
if isinstance(tensor, VkTensor) and tensor.constant_id < 0:
self.assertEqual(
tensor.storage_type, VkStorageType.TEXTURE_3D
)
self._run(edge, model, inputs, atol=0, rtol=0)

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

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

for shape, dim, supported in (
((7,), 0, True),
((0,), 0, True),
((3, 5), 0, True),
((3, 5), -1, True),
((0, 3), 0, True),
((2, 0, 3), 1, True),
((2, 0, 3), -1, True),
((1, 3, 5, 7), 1, True),
((1, 3, 5, 7), -2, True),
((2, 3, 5, 7), -1, True),
((2, 3, 5, 7), 2, True),
((1, 3, 1, 7), 1, True),
((1, 3, 5, 7), 0, False),
((2, 3, 5, 7), 1, False),
):
with self.subTest(shape=shape, dim=dim):
x = torch.zeros(shape, dtype=torch.bool)
if x.numel() > 0:
x[tuple(size // 2 for size in shape)] = True
model = AnyDim(dim)
edge = self._lower(model, (x,), fully_delegated=supported)
if not supported:
self.assertEqual(_vulkan_graphs(edge), [])
for graph in _vulkan_graphs(edge):
for value_id in graph.input_ids + graph.output_ids:
tensor = graph.values[value_id].value
self.assertEqual(tensor.storage_type, VkStorageType.TEXTURE_3D)
self._run(
edge,
model,
[(x,), (torch.zeros_like(x),), (torch.ones_like(x),)],
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),
(1, False),
(0, False),
):
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, 0, 16):
x = torch.zeros(2, s, 5, dtype=torch.bool)
if s not in (0, 3):
x[0, s // 2, 1] = True
x[1, -1, 3] = True
x[1, 0, 4] = True
inputs.append((x,))
seq = Dim("s", min=0, 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 @@ -706,7 +848,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 @@ -717,7 +859,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