Skip to content
Open
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
23 changes: 23 additions & 0 deletions backends/vulkan/test/test_vulkan_tensor_repr.py
Original file line number Diff line number Diff line change
Expand Up @@ -605,6 +605,29 @@ def test_unary_op_construction(self):
self.assertEqual(op_repsets.primary_arg_idx, 0)
self.assertTrue(op_repsets.sync_primary_io_repr)

def test_single_tensor_output_in_list_construction(self):
"""An op declared to return Tensor[] that yields exactly one tensor.

num_tensors_in_node() counts tensors rather than nesting, so such a
node reports 1 while meta["val"] is a one-element list rather than a
bare FakeTensor.
"""
arg = _make_tensor_arg_node((1, 3, 8, 8))
node = _make_op_node(
target=torch.ops.aten.split_with_sizes_copy.default,
args=(arg, [3]),
output_val=[_make_fake_tensor((1, 3, 8, 8))],
)

op_repsets = OpRepSets(
TensorRepSetList(ANY_STORAGE),
TensorRepSetList(ANY_STORAGE),
node,
DEFAULT_TEXTURE_LIMITS,
)

self.assertFalse(op_repsets.any_is_empty())

def test_binary_op_syncs_args(self):
"""When a single repset covers all inputs, sync_args_repr is True."""
op_repsets = self._make_binary_op()
Expand Down
9 changes: 8 additions & 1 deletion backends/vulkan/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -1442,8 +1442,15 @@ def __init__( # noqa: C901
outs_repset_list = TensorRepSetList([])
common_out_repset = ANY_STORAGE_INCL_PACKED_INT8
if num_tensors_in_node(op_node) == 1:
out_val = op_node.meta["val"]
# num_tensors_in_node counts tensors, not nesting: an op declared
# to return Tensor[] still lands here when it happens to produce
# exactly one, and meta["val"] is then a one-element list rather
# than a bare FakeTensor.
if isinstance(out_val, (list, tuple)):
out_val = out_val[0]
common_out_repset = filter_invalid_reprs(
op_node.meta["val"], outputs_repsets[0], texture_limits
out_val, outputs_repsets[0], texture_limits
)
outs_repset_list.append(common_out_repset)
# Multiple output tensors
Expand Down
Loading