PyTorch torch.repeat_interleave Function
PyTorch torch Reference Manual
torch.repeat_interleaveIt is a function in PyTorch used to repeat elements along a specified dimension.
Function Definition
torch.repeat_interleave(input, repeats, dim=None, *, output_size=None)
Usage Examples
Example
import torch
# Without specifying dim, repeat along all elements
x = torch.tensor([1, 2, 3])
result = torch.repeat_interleave(x, 2)
print("Each element repeated 2 times:")
print(result)
# Repeat along dim=0
y = torch.tensor([[1, 2], [3, 4]])
print("nOriginal tensor:")
print(y)
result = torch.repeat_interleave(y, 2, dim=0)
print("Repeated 2 times along dim=0:")
print(result)
# Different repetition counts for each element
result = torch.repeat_interleave(y, torch.tensor([1, 2]), dim=0)
print("nDifferent repetition counts along dim=0 [1, 2]:")
print(result)
# Repeat along dim=1
result = torch.repeat_interleave(y, 3, dim=1)
print("nRepeated 3 times along dim=1:")
print(result)
# Return output_size
result = torch.repeat_interleave(x, 2, output_size=9)
print("nSpecified output_size=9:")
print(result)
# Without specifying dim, repeat along all elements
x = torch.tensor([1, 2, 3])
result = torch.repeat_interleave(x, 2)
print("Each element repeated 2 times:")
print(result)
# Repeat along dim=0
y = torch.tensor([[1, 2], [3, 4]])
print("nOriginal tensor:")
print(y)
result = torch.repeat_interleave(y, 2, dim=0)
print("Repeated 2 times along dim=0:")
print(result)
# Different repetition counts for each element
result = torch.repeat_interleave(y, torch.tensor([1, 2]), dim=0)
print("nDifferent repetition counts along dim=0 [1, 2]:")
print(result)
# Repeat along dim=1
result = torch.repeat_interleave(y, 3, dim=1)
print("nRepeated 3 times along dim=1:")
print(result)
# Return output_size
result = torch.repeat_interleave(x, 2, output_size=9)
print("nSpecified output_size=9:")
print(result)
The output result is:
每个元素重复 2 次:
tensor([1, 1, 2, 2, 3, 3])
原始张量:
tensor([[1, 2],
[3, 4]])
沿 dim=0 重复 2 次:
tensor([[1, 2],
[1, 2],
[3, 4],
[3, 4]])
沿 dim=0 不同重复次数 [1, 2]:
tensor([[1, 2],
[3, 4],
[3, 4]])
沿 dim=1 重复 3 次:
tensor([[1, 1, 1, 2, 2, 2],
[3, 3, 3, 4, 4, 4]])
指定 output_size=9:
tensor([1, 1, 2, 2, 3, 3, 1, 2, 3])
Other Extensions