Skip to content
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