Skip to content
23 changes: 23 additions & 0 deletions backends/vulkan/runtime/VulkanBackend.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -703,6 +725,7 @@ class VulkanBackend final : public ::executorch::runtime::BackendInterface {
processed->Free();

if (err != Error::Ok) {
compute_graph->~ComputeGraph();
return err;
}

Expand Down
1 change: 1 addition & 0 deletions backends/vulkan/runtime/graph/ops/impl/UnaryOp.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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
3 changes: 2 additions & 1 deletion backends/vulkan/runtime/graph/ops/utils/StagingUtils.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
63 changes: 63 additions & 0 deletions backends/vulkan/test/test_vulkan_dynamic.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,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,
Expand Down Expand Up @@ -166,6 +175,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):
Expand Down Expand Up @@ -419,6 +461,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()
3 changes: 2 additions & 1 deletion backends/vulkan/test/utils/test_utils.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
Loading