Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
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 @@ -1602,12 +1602,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 @@ -32,6 +32,8 @@

from executorch.exir import EdgeCompileConfig, to_edge_transform_and_lower

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 @@ -160,6 +162,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 @@ -172,6 +212,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 @@ -205,6 +261,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 @@ -231,6 +303,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