From 8d46517fba561ebdad848a03a6709ea295979c81 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 | 2 +- backends/vulkan/test/test_vulkan_dynamic.py | 69 +++++++++++++++++++++ 2 files changed, 70 insertions(+), 1 deletion(-) diff --git a/backends/vulkan/op_registry.py b/backends/vulkan/op_registry.py index 1bdff20f05c..efe9ca7d81f 100644 --- a/backends/vulkan/op_registry.py +++ b/backends/vulkan/op_registry.py @@ -1367,7 +1367,7 @@ def register_expand_copy(): return OpFeatures( inputs_storage=utils.ANY_STORAGE, inputs_dtypes=utils.FP_INT_BOOL_T, - supports_resize=False, + supports_resize=True, supports_highdim=True, ) diff --git a/backends/vulkan/test/test_vulkan_dynamic.py b/backends/vulkan/test/test_vulkan_dynamic.py index c9292d916c1..7b760cd5b2f 100644 --- a/backends/vulkan/test/test_vulkan_dynamic.py +++ b/backends/vulkan/test/test_vulkan_dynamic.py @@ -50,6 +50,25 @@ def _vulkan_graphs(edge): ] +class TransformerBlock(torch.nn.Module): + def __init__(self, sdpa): + super().__init__() + self.sdpa = sdpa + self.qkv = torch.nn.Linear(64, 192) + self.ff = torch.nn.Linear(64, 64) + + def forward(self, x, lengths): + b, s, _ = x.shape + mask = (torch.arange(s)[None, :] < lengths[:, None])[:, None, None, :] + q, k, v = self.qkv(x).view(b, s, 3, 2, 32).permute(2, 0, 3, 1, 4) + if self.sdpa: + y = F.scaled_dot_product_attention(q, k, v, attn_mask=mask) + else: + bias = torch.where(mask, 0.0, -torch.inf) + y = torch.softmax(q @ k.transpose(-1, -2) * 32**-0.5 + bias, -1) @ v + return F.gelu(self.ff(y.transpose(1, 2).reshape(b, s, 64))) + + class ConstantMask(torch.nn.Module): def __init__(self): super().__init__() @@ -144,6 +163,36 @@ def _run( ) ) + def test_partition_transformer(self): + for sdpa in (False, True): + with self.subTest(sdpa=sdpa): + self._lower( + TransformerBlock(sdpa), + (torch.randn(1, 16, 64), torch.tensor([16])), + ({1: Dim("s", min=2, max=1000)}, {}), + ) + + def test_dynamic_transformer(self): + torch.manual_seed(0) + for sdpa in (False, True): + for storage in (VkStorageType.TEXTURE_3D, VkStorageType.BUFFER): + with self.subTest(sdpa=sdpa, storage=storage): + model = TransformerBlock(sdpa).eval() + inputs = [ + (torch.randn(1, s, 64), torch.tensor([length])) + for s, length in ( + (16, 16), + (7, 4), + (31, 23), + (2, 1), + (16, 0 if sdpa else 5), + ) + ] + edge = self._lower( + model, inputs[0], ({1: Dim("s", min=2, max=32)}, {}), storage + ) + self._run(edge, model, inputs) + def test_partition_any_unsupported_inputs(self): class AnyDim(torch.nn.Module): def forward(self, x): @@ -249,6 +298,26 @@ def forward(self, x): edge = self._lower(model, inputs[0]) self._run(edge, model, inputs, atol=0, rtol=0) + def test_dynamic_expand(self): + class Expand(torch.nn.Module): + def __init__(self): + super().__init__() + self.register_buffer( + "offset", torch.arange(4, dtype=torch.float32)[None, :] + ) + + def forward(self, x): + return x.expand(2, x.shape[1], 4) + self.offset.expand(2, 4)[:, None, :] + + for storage in (VkStorageType.TEXTURE_3D, VkStorageType.BUFFER): + with self.subTest(storage=storage): + model = Expand() + inputs = [(torch.randn(1, s, 1),) for s in (16, 3, 31, 2, 16)] + edge = self._lower( + model, inputs[0], ({1: Dim("s", min=2, max=32)},), storage + ) + self._run(edge, model, inputs, atol=0, rtol=0) + def test_dynamic_full(self): class Full(torch.nn.Module): def forward(self, x):