From fcf87f11daa0bdf7f226c126416faf662bb2a360 Mon Sep 17 00:00:00 2001 From: Mergen Nachin Date: Tue, 29 Sep 2026 11:02:07 -0400 Subject: [PATCH] [INITIAL] Recreate the Vulkan transformer and conformance stack with ghstack [ghstack-poisoned] --- backends/vulkan/op_registry.py | 14 ++++ .../vulkan/runtime/graph/ops/glsl/reduce.glsl | 15 +++-- .../vulkan/runtime/graph/ops/glsl/reduce.yaml | 3 + .../graph/ops/glsl/reduce_per_row_buffer.yaml | 4 ++ .../vulkan/runtime/graph/ops/impl/Reduce.cpp | 12 ++++ backends/vulkan/test/test_vulkan_dynamic.py | 64 ++++++++++++++++++- 6 files changed, 103 insertions(+), 9 deletions(-) diff --git a/backends/vulkan/op_registry.py b/backends/vulkan/op_registry.py index 6d672a6ceb4..1bdff20f05c 100644 --- a/backends/vulkan/op_registry.py +++ b/backends/vulkan/op_registry.py @@ -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 # ============================================================================= diff --git a/backends/vulkan/runtime/graph/ops/glsl/reduce.glsl b/backends/vulkan/runtime/graph/ops/glsl/reduce.glsl index f37a0fc773a..cb6cd8423d6 100644 --- a/backends/vulkan/runtime/graph/ops/glsl/reduce.glsl +++ b/backends/vulkan/runtime/graph/ops/glsl/reduce.glsl @@ -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)} @@ -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" @@ -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)); @@ -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 @@ -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]); } @@ -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))); } } diff --git a/backends/vulkan/runtime/graph/ops/glsl/reduce.yaml b/backends/vulkan/runtime/graph/ops/glsl/reduce.yaml index f85d8405b2d..550aeea093a 100644 --- a/backends/vulkan/runtime/graph/ops/glsl/reduce.yaml +++ b/backends/vulkan/runtime/graph/ops/glsl/reduce.yaml @@ -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 diff --git a/backends/vulkan/runtime/graph/ops/glsl/reduce_per_row_buffer.yaml b/backends/vulkan/runtime/graph/ops/glsl/reduce_per_row_buffer.yaml index e5a94165b96..e4850fb1a44 100644 --- a/backends/vulkan/runtime/graph/ops/glsl/reduce_per_row_buffer.yaml +++ b/backends/vulkan/runtime/graph/ops/glsl/reduce_per_row_buffer.yaml @@ -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 diff --git a/backends/vulkan/runtime/graph/ops/impl/Reduce.cpp b/backends/vulkan/runtime/graph/ops/impl/Reduce.cpp index 23c8f000bed..2c29045907c 100644 --- a/backends/vulkan/runtime/graph/ops/impl/Reduce.cpp +++ b/backends/vulkan/runtime/graph/ops/impl/Reduce.cpp @@ -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& args) { + if (graph.is_buffer_storage(args[0])) { + VK_CHECK_COND( + normalize( + graph.extract_scalar(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 diff --git a/backends/vulkan/test/test_vulkan_dynamic.py b/backends/vulkan/test/test_vulkan_dynamic.py index af75c099112..c9292d916c1 100644 --- a/backends/vulkan/test/test_vulkan_dynamic.py +++ b/backends/vulkan/test/test_vulkan_dynamic.py @@ -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): @@ -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): @@ -683,7 +739,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), @@ -694,7 +750,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: