Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
58 commits
Select commit Hold shift + click to select a range
aae8159
[INITIAL] Recreate the Vulkan transformer and conformance stack with …
mergennachin Sep 29, 2026
bcab9d0
[INITIAL] Recreate the Vulkan transformer and conformance stack with …
mergennachin Sep 29, 2026
d6db604
[INITIAL] Recreate the Vulkan transformer and conformance stack with …
mergennachin Sep 29, 2026
446a5d6
[INITIAL] Recreate the Vulkan transformer and conformance stack with …
mergennachin Sep 29, 2026
fb905cf
[INITIAL] Recreate the Vulkan transformer and conformance stack with …
mergennachin Sep 29, 2026
2131a32
[INITIAL] Recreate the Vulkan transformer and conformance stack with …
mergennachin Sep 29, 2026
f310b91
[INITIAL] Recreate the Vulkan transformer and conformance stack with …
mergennachin Sep 29, 2026
f588b1e
[INITIAL] Recreate the Vulkan transformer and conformance stack with …
mergennachin Sep 29, 2026
836dc54
[INITIAL] Recreate the Vulkan transformer and conformance stack with …
mergennachin Sep 29, 2026
b708c1a
[INITIAL] Recreate the Vulkan transformer and conformance stack with …
mergennachin Sep 29, 2026
0f3e873
[UPDATE] Rebase remaining Vulkan stack onto main after parts 1–3 landed
mergennachin Oct 3, 2026
ee8dffb
[UPDATE] Rebase remaining Vulkan stack onto main after parts 1–3 landed
mergennachin Oct 3, 2026
e1b9cb7
[UPDATE] Rebase remaining Vulkan stack onto main after parts 1–3 landed
mergennachin Oct 3, 2026
f089d4f
[UPDATE] Rebase remaining Vulkan stack onto main after parts 1–3 landed
mergennachin Oct 3, 2026
ea19b69
[UPDATE] Rebase remaining Vulkan stack onto main after parts 1–3 landed
mergennachin Oct 3, 2026
b066808
[UPDATE] Rebase remaining Vulkan stack onto main after parts 1–3 landed
mergennachin Oct 3, 2026
a53613d
[UPDATE] Rebase remaining Vulkan stack onto main after parts 1–3 landed
mergennachin Oct 3, 2026
f963896
[UPDATE] Cover exact FP16 halfway rounding in texture reductions
mergennachin Oct 3, 2026
5b423ab
[UPDATE] Cover exact FP16 halfway rounding in texture reductions
mergennachin Oct 3, 2026
5dd4a73
[UPDATE] Cover exact FP16 halfway rounding in texture reductions
mergennachin Oct 3, 2026
a6a7c39
[UPDATE] Cover exact FP16 halfway rounding in texture reductions
mergennachin Oct 3, 2026
6f39097
[UPDATE] Cover exact FP16 halfway rounding in texture reductions
mergennachin Oct 3, 2026
ccc3c43
[UPDATE] Cover exact FP16 halfway rounding in texture reductions
mergennachin Oct 3, 2026
a231002
[UPDATE] Match arg-reduction partitioning to the buffer-only last-dim…
mergennachin Oct 3, 2026
cf96624
[UPDATE] Match arg-reduction partitioning to the buffer-only last-dim…
mergennachin Oct 3, 2026
e066a11
[UPDATE] Match arg-reduction partitioning to the buffer-only last-dim…
mergennachin Oct 3, 2026
b378bbd
[UPDATE] Match arg-reduction partitioning to the buffer-only last-dim…
mergennachin Oct 3, 2026
83d3ac3
[UPDATE] Match arg-reduction partitioning to the buffer-only last-dim…
mergennachin Oct 3, 2026
18ae80a
[UPDATE] Match arg-reduction partitioning to the buffer-only last-dim…
mergennachin Oct 3, 2026
7baf82c
[UPDATE] Add exact FP16 zero/subnormal and overflow-threshold coverage
mergennachin Oct 3, 2026
213b59f
[UPDATE] Add exact FP16 zero/subnormal and overflow-threshold coverage
mergennachin Oct 3, 2026
707344f
[UPDATE] Add exact FP16 zero/subnormal and overflow-threshold coverage
mergennachin Oct 3, 2026
2c83efe
[UPDATE] Add exact FP16 zero/subnormal and overflow-threshold coverage
mergennachin Oct 3, 2026
fb9b582
[UPDATE] Add exact FP16 zero/subnormal and overflow-threshold coverage
mergennachin Oct 3, 2026
a1453d9
[UPDATE] Add exact FP16 zero/subnormal and overflow-threshold coverage
mergennachin Oct 3, 2026
1ad69df
[UPDATE] Exercise int32 reduction shaders directly beyond the removed…
mergennachin Oct 3, 2026
59e4479
[UPDATE] Exercise int32 reduction shaders directly beyond the removed…
mergennachin Oct 3, 2026
c5e36bc
[UPDATE] Exercise int32 reduction shaders directly beyond the removed…
mergennachin Oct 3, 2026
8d43512
[UPDATE] Exercise int32 reduction shaders directly beyond the removed…
mergennachin Oct 3, 2026
6942500
[UPDATE] Exercise int32 reduction shaders directly beyond the removed…
mergennachin Oct 3, 2026
554ab97
[UPDATE] Exercise int32 reduction shaders directly beyond the removed…
mergennachin Oct 3, 2026
b8177d5
[UPDATE] Support any.dim without keepdim using texture reduction and …
mergennachin Oct 3, 2026
c312284
[UPDATE] Support any.dim without keepdim using texture reduction and …
mergennachin Oct 3, 2026
019f9be
[UPDATE] Support any.dim without keepdim using texture reduction and …
mergennachin Oct 3, 2026
0b3b02f
[UPDATE] Support any.dim without keepdim using texture reduction and …
mergennachin Oct 3, 2026
6c1b29c
[UPDATE] Support any.dim without keepdim using texture reduction and …
mergennachin Oct 3, 2026
bcd89c3
[UPDATE] Rebase remaining Vulkan stack after #23245 landed
mergennachin Oct 5, 2026
e75b412
[UPDATE] Rebase remaining Vulkan stack after #23245 landed
mergennachin Oct 5, 2026
a6ea4e1
[UPDATE] Rebase remaining Vulkan stack after #23245 landed
mergennachin Oct 5, 2026
1b7b3fd
[UPDATE] Rebase remaining Vulkan stack after #23245 landed
mergennachin Oct 5, 2026
0e50614
[UPDATE] Update
mergennachin Oct 7, 2026
2d8dbfe
[UPDATE] Update
mergennachin Oct 7, 2026
555d25c
[UPDATE] Update
mergennachin Oct 7, 2026
a137634
[UPDATE] Address SS-JIA review on full fill UBO types and rebase afte…
mergennachin Oct 7, 2026
eccbf5f
[UPDATE] Address SS-JIA review on full fill UBO types and rebase afte…
mergennachin Oct 7, 2026
80283ae
[UPDATE] Rebase remaining Vulkan stack onto main ad434a39ea
mergennachin Oct 7, 2026
d248f3a
[UPDATE] Rebase remaining Vulkan stack onto main ad434a39ea
mergennachin Oct 7, 2026
a7e71de
[UPDATE] Rebase remaining Vulkan stack after #23248 landed
mergennachin Oct 7, 2026
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
11 changes: 10 additions & 1 deletion backends/vulkan/op_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -1610,12 +1610,21 @@ def register_full_cpp_ops():
# =============================================================================


@update_features(exir_ops.edge.aten.scalar_tensor.default)
@update_features(
[
exir_ops.edge.aten.scalar_tensor.default,
# EXIR deliberately keeps scalar_tensor in the ATen dialect.
torch.ops.aten.scalar_tensor.default,
]
)
def register_scalar_tensor():
return OpFeatures(
inputs_storage=utils.CHANNELS_PACKED_TEXTURE,
inputs_dtypes=utils.FP_INT_T,
supports_resize=True,
are_node_inputs_supported_fn=lambda node: is_scalar_value_supported(
node.args[0], node.meta["val"].dtype
),
)


Expand Down
3 changes: 3 additions & 0 deletions backends/vulkan/runtime/graph/ops/glsl/scalar_tensor.glsl
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@ ${define_explicit_type_extensions(SCALAR_VALUE_TYPE)}
${define_active_storage_type(STORAGE)}

#include "indexing_utils.h"
#include "convert.glslh"

layout(std430) buffer;

Expand Down Expand Up @@ -52,6 +53,8 @@ void main() {
}

VEC4_T outtex = VEC4_T(scalar_value);
$if DTYPE == "half":
outtex = round_to_half_rte(outtex);
write_texel(t_out, pos, outtex);
}

Expand Down
20 changes: 9 additions & 11 deletions backends/vulkan/runtime/graph/ops/glsl/scalar_tensor.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -12,16 +12,14 @@ scalar_tensor:
PACKING: C_packed
STORAGE: texture3d
generate_variant_forall:
DTYPE:
- VALUE: half
- VALUE: float
- VALUE: int32
STORAGE:
- VALUE: texture3d
- VALUE: buffer
SCALAR_VALUE_TYPE:
- VALUE: float
- VALUE: int32
- VALUE: bool
combination:
parameter_names: [DTYPE, STORAGE, SCALAR_VALUE_TYPE]
combos:
- parameter_values: [half, texture3d, float]
- parameter_values: [half, buffer, float]
- parameter_values: [float, texture3d, float]
- parameter_values: [float, buffer, float]
- parameter_values: [int32, texture3d, int32]
- parameter_values: [int32, buffer, int32]
shader_variants:
- NAME: scalar_tensor
10 changes: 7 additions & 3 deletions backends/vulkan/runtime/graph/ops/impl/ScalarTensor.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -16,17 +16,21 @@ namespace vkcompute {
void scalar_tensor(ComputeGraph& graph, const std::vector<ValueRef>& args) {
// Extract the scalar value from the first argument
ValueRef scalar_in = args[0];
float scalar_value = graph.extract_scalar<float>(scalar_in);

// Get the output tensor reference
ValueRef out = args[args.size() - 1];
const vkapi::ScalarType scalar_dtype =
graph.dtype_of(out) == vkapi::kInt ? vkapi::kInt : vkapi::kFloat;
const vkapi::BufferBindInfo scalar_buffer = scalar_dtype == vkapi::kInt
? graph.create_params_buffer(graph.extract_scalar<int32_t>(scalar_in))
: graph.create_params_buffer(graph.extract_scalar<float>(scalar_in));

std::string kernel_name("scalar_tensor");
kernel_name.reserve(kShaderNameReserve);

add_dtype_suffix(kernel_name, graph.dtype_of(out));
add_storage_type_suffix(kernel_name, graph.storage_type_of(out));
add_dtype_suffix(kernel_name, graph.dtype_of(scalar_in));
add_dtype_suffix(kernel_name, scalar_dtype);

graph.execute_nodes().emplace_back(new DispatchNode(
graph,
Expand All @@ -36,7 +40,7 @@ void scalar_tensor(ComputeGraph& graph, const std::vector<ValueRef>& args) {
// Inputs and Outputs
{{out, vkapi::kWrite}},
// Shader params buffers
{graph.create_params_buffer(scalar_value)},
{scalar_buffer},
// Push Constants
{},
// Specialization Constants
Expand Down
5 changes: 4 additions & 1 deletion backends/vulkan/serialization/vulkan_graph_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -474,10 +474,13 @@ def process_call_function_node(self, node) -> None:
if not self.delegate_mapping_builder
else self.delegate_mapping_builder.insert_delegate_mapping_entry(node)
)
operator_name = node.target.__name__
if node.target == torch.ops.aten.scalar_tensor.default:
operator_name = "aten.scalar_tensor.default"
self.chain.append(
vk_graph_schema.OperatorCall(
node_id=operator_node_id, # pyre-ignore[6]: this is going to be an int
name=node.target.__name__,
name=operator_name,
args=operator_call_args,
),
)
Expand Down
2 changes: 2 additions & 0 deletions backends/vulkan/test/op_tests/cases.py
Original file line number Diff line number Diff line change
Expand Up @@ -892,10 +892,12 @@ def get_scalar_tensor_inputs():
test_suite = VkTestSuite(
[
(42.0,),
(42,),
(3.14,),
(2.72,),
(0.0,),
(-1.0,),
(-7,),
(100.0,),
]
)
Expand Down
102 changes: 102 additions & 0 deletions backends/vulkan/test/test_vulkan_dynamic.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,8 @@

from executorch.exir.backend.backend_api import to_backend

from executorch.exir.dialects._ops import ops as exir_ops

from executorch.exir.lowered_backend_module import LoweredBackendModule

from torch.export import Dim, export
Expand Down Expand Up @@ -165,6 +167,44 @@ def test_dynamic_gelu(self):
tolerance = 5e-6 if dtype == torch.float32 else 1e-3
self._run(edge, model, inputs, atol=tolerance, rtol=tolerance)

def test_scalar_tensor_values(self):
class WhereScalars(torch.nn.Module):
def __init__(self, positive, negative):
super().__init__()
self.positive = positive
self.negative = negative

def forward(self, x):
return torch.where(x, self.positive, self.negative)

inputs = [(torch.tensor([True, False, True, False]),)]
for positive, negative in (
(3, -7.0),
(3.0, -7),
(16777217, -7),
(2**31 - 1, -(2**31)),
):
with self.subTest(positive=positive, negative=negative):
model = WhereScalars(positive, negative)
fully_delegated = isinstance(positive, float) or isinstance(
negative, float
)
edge = self._lower(model, inputs[0], fully_delegated=fully_delegated)
if not fully_delegated:
self.assertEqual(
[
node.target
for node in edge.exported_program().graph.nodes
if node.op == "call_function"
and node.target != operator.getitem
],
[
torch.ops.higher_order.executorch_call_delegate,
exir_ops.edge.aten.where.self,
],
)
self._run(edge, model, inputs, atol=0, rtol=0)

def test_gelu_with_singleton_dimensions(self):
for approximate in ("none", "tanh"):
for shape in ((6, 1, 3), (2, 1, 3, 5)):
Expand All @@ -177,6 +217,22 @@ 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_scalar_tensor_dtypes(self):
class ScalarTensor(torch.nn.Module):
def __init__(self, dtype):
super().__init__()
self.dtype = dtype

def forward(self, x):
return torch.scalar_tensor(2.5, dtype=self.dtype)

inputs = [(torch.ones(1),)]
for dtype in (torch.float16, torch.float32, torch.int32):
with self.subTest(dtype=dtype):
model = ScalarTensor(dtype)
edge = self._lower(model, inputs[0])
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 @@ -210,6 +266,22 @@ def test_constant_bool_mask(self):
)
self._run(edge, model, inputs, atol=0, rtol=0)

def test_scalar_types_before_conv_and_view(self):
class ScalarTypes(torch.nn.Module):
def __init__(self):
super().__init__()
self.conv = torch.nn.Conv2d(1, 2, 3, padding=1)

def forward(self, x):
y = self.conv(torch.where(x > 0, 1.0, 0.5))
return y.view(1, 2, x.shape[2], -1)

torch.manual_seed(0)
model = ScalarTypes().eval()
inputs = [(torch.randn(1, 1, s, 5),) for s in (7, 2, 15, 7)]
edge = self._lower(model, inputs[0], ({2: Dim("s", min=2, max=16)},))
self._run(edge, model, inputs)

def test_nan_scalars_fall_back(self):
class NanScalar(torch.nn.Module):
def __init__(self, kind):
Expand All @@ -236,6 +308,36 @@ def forward(self, x):
edge = self._lower(model, inputs[0], fully_delegated=False)
self._run(edge, model, inputs, atol=0, rtol=0, equal_nan=True)

def test_integer_scalar_range_fallback(self):
class LargeScalar(torch.nn.Module):
def __init__(self, kind, value):
super().__init__()
self.kind = kind
self.value = value

def forward(self, x):
if self.kind == "scalar_tensor":
return torch.scalar_tensor(self.value, dtype=torch.int64)
if self.kind == "full":
return torch.full(x.shape, self.value, dtype=torch.int64)
return torch.full_like(x, self.value, dtype=torch.int64)

inputs = [(torch.zeros(2, 3),)]
for kind in ("scalar_tensor", "full", "full_like"):
for value in (
2**31 - 0.5,
-(2**31) - 0.5,
2**31,
2**40,
2**63 - 1,
-(2**63),
):
with self.subTest(kind=kind, value=value):
model = LargeScalar(kind, value)
edge = self._lower(model, inputs[0], fully_delegated=False)
self.assertEqual(_vulkan_graphs(edge), [])
self._run(edge, model, inputs, atol=0, rtol=0)

def test_64_bit_arithmetic_without_downcasting(self):
class Arithmetic(torch.nn.Module):
def forward(self, x):
Expand Down
22 changes: 22 additions & 0 deletions backends/vulkan/test/test_vulkan_graph_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,28 @@ def test_scalar_cache_preserves_types_and_signed_zero(self):
self.assertEqual(builder.get_or_create_scalar_value(value), value_id)
self.assertEqual(repr(builder.values[value_id].value), repr(serialized))

def test_aten_scalar_tensor_keeps_namespace(self):
class Mask(torch.nn.Module):
def forward(self, x):
return torch.where(x, 0.0, -torch.inf)

program = torch.export.export(Mask(), (torch.tensor([True, False]),))
edge = to_edge(program)
program = apply_passes(edge.exported_program(), [SpecPropPass()])
self.assertEqual(
sum(
node.target == torch.ops.aten.scalar_tensor.default
for node in program.graph.nodes
),
2,
)
graph = VkGraphBuilder(
program, DelegateMappingBuilder(generated_identifiers=True)
).build_graph()
names = [op.name for op in graph.chain]
self.assertEqual(names.count("aten.scalar_tensor.default"), 2)
self.assertNotIn("scalar_tensor.default", names)


class TestVkGraphBuilderInputIds(unittest.TestCase):
"""The serialized input list has to match the delegate call's arguments.
Expand Down
Loading