diff --git a/backends/vulkan/runtime/VulkanBackend.cpp b/backends/vulkan/runtime/VulkanBackend.cpp index b7d3de66f01..3a11adf05e0 100644 --- a/backends/vulkan/runtime/VulkanBackend.cpp +++ b/backends/vulkan/runtime/VulkanBackend.cpp @@ -660,6 +660,28 @@ class VulkanBackend final : public ::executorch::runtime::BackendInterface { VkGraphPtr flatbuffer_graph = vkgraph::GetVkGraph(flatbuffer_data); + if (!compute_graph->context() + ->adapter_ptr() + ->supports_8bit_storage_buffers()) { + for (const auto* value : *flatbuffer_graph->values()) { + const auto* tensor = value->value_as_VkTensor(); + // Constants become CPU TensorRefs; their prepack destinations are + // separate GPU tensors checked here. + if (tensor == nullptr || tensor->constant_id() >= 0 || + tensor->datatype() != vkgraph::VkDataType::BOOL) { + continue; + } + const auto storage = + tensor->storage_type() == vkgraph::VkStorageType::DEFAULT_STORAGE + ? compute_graph->suggested_storage_type() + : get_storage_type(tensor->storage_type()); + ET_CHECK_OR_RETURN_ERROR( + storage != utils::kBuffer, + NotSupported, + "Vulkan bool buffer tensors require 8-bit storage buffer support"); + } + } + GraphBuilder builder( compute_graph, flatbuffer_graph, @@ -703,6 +725,7 @@ class VulkanBackend final : public ::executorch::runtime::BackendInterface { processed->Free(); if (err != Error::Ok) { + compute_graph->~ComputeGraph(); return err; } diff --git a/backends/vulkan/runtime/graph/ops/impl/UnaryOp.cpp b/backends/vulkan/runtime/graph/ops/impl/UnaryOp.cpp index d17979a4fa2..af0711f3d51 100644 --- a/backends/vulkan/runtime/graph/ops/impl/UnaryOp.cpp +++ b/backends/vulkan/runtime/graph/ops/impl/UnaryOp.cpp @@ -231,6 +231,7 @@ REGISTER_OPERATORS { VK_REGISTER_OP(aten.log10.default, log10); VK_REGISTER_OP(aten.round.default, round); VK_REGISTER_OP(aten.bitwise_not.default, bitwise_not); + VK_REGISTER_OP(aten.logical_not.default, bitwise_not); } } // namespace vkcompute diff --git a/backends/vulkan/runtime/graph/ops/utils/StagingUtils.cpp b/backends/vulkan/runtime/graph/ops/utils/StagingUtils.cpp index 837e509e691..ca52ba3cb60 100644 --- a/backends/vulkan/runtime/graph/ops/utils/StagingUtils.cpp +++ b/backends/vulkan/runtime/graph/ops/utils/StagingUtils.cpp @@ -17,7 +17,8 @@ namespace vkcompute { bool is_bitw8(vkapi::ScalarType dtype) { return dtype == vkapi::kByte || dtype == vkapi::kChar || - dtype == vkapi::kQInt8 || dtype == vkapi::kQUInt8; + dtype == vkapi::kQInt8 || dtype == vkapi::kQUInt8 || + dtype == vkapi::kBool; } vkapi::ShaderInfo get_nchw_to_tensor_shader( diff --git a/backends/vulkan/test/test_vulkan_dynamic.py b/backends/vulkan/test/test_vulkan_dynamic.py index cfb8de62f54..6feb41c65a5 100644 --- a/backends/vulkan/test/test_vulkan_dynamic.py +++ b/backends/vulkan/test/test_vulkan_dynamic.py @@ -46,6 +46,15 @@ def _vulkan_graphs(edge): ] +class ConstantMask(torch.nn.Module): + def __init__(self): + super().__init__() + self.register_buffer("mask", torch.arange(21).reshape(3, 7) % 2 == 0) + + def forward(self, x): + return torch.where(self.mask, x, -x) + + class TestVulkanDynamic(unittest.TestCase): def _lower( self, @@ -161,6 +170,39 @@ 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_dynamic_logical_not(self): + class LogicalNot(torch.nn.Module): + def forward(self, x): + return torch.logical_not(x) + + model = LogicalNot() + inputs = [ + ((torch.arange(3 * s).reshape(3, s) % 3 == 0),) for s in (7, 2, 15, 3, 7) + ] + for storage in (VkStorageType.TEXTURE_3D, VkStorageType.BUFFER): + with self.subTest(storage=storage): + edge = self._lower( + model, inputs[0], ({1: Dim("s", min=2, max=16)},), storage + ) + self._run(edge, model, inputs, atol=0, rtol=0) + + def test_constant_bool_mask(self): + model = ConstantMask() + inputs = [(torch.linspace(-1, 1, 21).reshape(3, 7),)] + for storage in (VkStorageType.TEXTURE_3D, VkStorageType.BUFFER): + with self.subTest(storage=storage): + edge = self._lower(model, inputs[0], storage=storage) + self.assertTrue( + any( + isinstance(value.value, VkTensor) + and value.value.constant_id >= 0 + and value.value.datatype == VkDataType.BOOL + for graph in _vulkan_graphs(edge) + for value in graph.values + ) + ) + self._run(edge, model, inputs, atol=0, rtol=0) + def test_4d_reductions(self): class Reduce(torch.nn.Module): def __init__(self, op, dim): @@ -289,6 +331,27 @@ def forward(self, x): edge = self._lower(model, (x,), storage=VkStorageType.BUFFER) self._run(edge, model, [(x,)], atol=0, rtol=0) + @unittest.skipUnless(USING_SWIFTSHADER, "requires a device without 8-bit buffers") + def test_bool_buffers_fail_cleanly_without_8bit_storage(self): + from executorch.extension.pybindings.portable_lib import ( + _load_for_executorch_from_buffer, + ) + + class LogicalNot(torch.nn.Module): + def forward(self, x): + return torch.logical_not(x) + + for model, inputs in ( + (LogicalNot(), (torch.zeros(3, 7, dtype=torch.bool),)), + (ConstantMask(), (torch.zeros(3, 7),)), + ): + with self.subTest(model=type(model).__name__): + edge = self._lower(model, inputs, storage=VkStorageType.BUFFER) + program_buffer = edge.to_executorch().buffer + module = _load_for_executorch_from_buffer(program_buffer) + with self.assertRaisesRegex(RuntimeError, r"0x:?10\b"): + module.run_method("forward", inputs) + if __name__ == "__main__": unittest.main() diff --git a/backends/vulkan/test/utils/test_utils.cpp b/backends/vulkan/test/utils/test_utils.cpp index fdd45baa6b7..31d7fe739b9 100644 --- a/backends/vulkan/test/utils/test_utils.cpp +++ b/backends/vulkan/test/utils/test_utils.cpp @@ -35,7 +35,8 @@ GlobalWorkGrid make_linear_dispatch( bool is_bitw8(vkapi::ScalarType dtype) { return dtype == vkapi::kByte || dtype == vkapi::kChar || - dtype == vkapi::kQInt8 || dtype == vkapi::kQUInt8; + dtype == vkapi::kQInt8 || dtype == vkapi::kQUInt8 || + dtype == vkapi::kBool; } vkapi::ShaderInfo get_nchw_to_tensor_shader(