Skip to content

Commit

Permalink
fix: allow torch backend split to infer final split size (#28230)
Browse files Browse the repository at this point in the history
  • Loading branch information
Sam-Armstrong authored Feb 9, 2024
1 parent d38a312 commit c60b3a2
Showing 1 changed file with 4 additions and 0 deletions.
4 changes: 4 additions & 0 deletions ivy/functional/backends/torch/manipulation.py
Original file line number Diff line number Diff line change
Expand Up @@ -241,6 +241,10 @@ def split(
torch.tensor(dim_size) / torch.tensor(num_or_size_splits)
)
elif isinstance(num_or_size_splits, list):
if num_or_size_splits[-1] == -1:
# infer the final size split
remaining_size = dim_size - sum(num_or_size_splits[:-1])
num_or_size_splits[-1] = remaining_size
num_or_size_splits = tuple(num_or_size_splits)
return list(torch.split(x, num_or_size_splits, axis))

Expand Down

0 comments on commit c60b3a2

Please sign in to comment.