PyTorch torch.take_along_dim Function
Pytorch torch Reference Manual
torch.take_along_dimIt is a function in PyTorch used to obtain elements at index positions along a specified dimension. It takes values fromindicesthe indices specified indimalong the dimension frominput.
Function Definition
torch.take_along_dim(input, indices, dim)
Parameters:
input(Tensor): Input tensor.indices(Tensor): Index tensor, specifying the positions of elements to extract. The shape must be compatible with input along the dim dimension.dim(int): The dimension along which to operate.
Return Value:
torch.Tensor: Returns a new tensor composed of the elements extracted by the indices.
Usage Example
Example
import torch
# Create a 2D tensor
x = torch.tensor([[1, 2, 3],
[4, 5, 6],
[7, 8, 9]])
print("Original tensor:")
print(x)
# Extract elements along dim=1
indices = torch.tensor([[0, 1, 2],
[2, 1, 0],
[0, 0, 0]])
y = torch.take_along_dim(x, indices, dim=1)
print("nIndices:")
print(indices)
print("nElements extracted along dim=1:")
print(y)
# Create a 2D tensor
x = torch.tensor([[1, 2, 3],
[4, 5, 6],
[7, 8, 9]])
print("Original tensor:")
print(x)
# Extract elements along dim=1
indices = torch.tensor([[0, 1, 2],
[2, 1, 0],
[0, 0, 0]])
y = torch.take_along_dim(x, indices, dim=1)
print("nIndices:")
print(indices)
print("nElements extracted along dim=1:")
print(y)
The output result is:
原始张量:
tensor([[1, 2, 3],
[4, 5, 6],
[7, 8, 9]])
索引:
tensor([[0, 1, 2],
[2, 1, 0],
[0, 0, 0]])
沿 dim=1 取出的元素:
tensor([[1, 2, 3],
[6, 5, 4],
[7, 7, 7]])
Example
import torch
# Extract elements along dim=0
x = torch.tensor([[1, 2, 3],
[4, 5, 6],
[7, 8, 9]])
# Take different rows for each column
indices = torch.tensor([[0, 1, 2],
[2, 0, 1],
[1, 2, 0]])
y = torch.take_along_dim(x, indices, dim=0)
print("Original tensor:")
print(x)
print("nIndices:")
print(indices)
print("nElements extracted along dim=0:")
print(y)
# Extract elements along dim=0
x = torch.tensor([[1, 2, 3],
[4, 5, 6],
[7, 8, 9]])
# Take different rows for each column
indices = torch.tensor([[0, 1, 2],
[2, 0, 1],
[1, 2, 0]])
y = torch.take_along_dim(x, indices, dim=0)
print("Original tensor:")
print(x)
print("nIndices:")
print(indices)
print("nElements extracted along dim=0:")
print(y)
The output result is:
原始张量:
tensor([[1, 2, 3],
[4, 5, 6],
[7, 8, 9]])
索引:
tensor([[0, 1, 2],
[2, 0, 1],
[1, 2, 0]])
沿 dim=0 取出的元素:
tensor([[1, 5, 9],
[7, 2, 6],
[4, 8, 3]])
</p>
<div class="example">
<h2 class="example">实例</h2>
<div class="example_code">
<span style="color: Green;font-weight:bold;">import</span> torch<br />
<br />
<span style="color: #a50"># 在3D张量上使用</span><br />
x <span style="color: Gray;">=</span> torch.<span style="color: #05a;">arange</span><span style="color: Olive;">(</span><span style="color: Maroon;">24</span><span style="color: Olive;">)</span>.<span style="color: #05a;">reshape</span><span style="color: Olive;">(</span><span style="color: Maroon;">2</span><span style="color: Gray;">,</span> <span style="color: Maroon;">3</span><span style="color: Gray;">,</span> <span style="color: Maroon;">4</span><span style="color: Olive;">)</span><br />
<span style="color: Green;font-weight:bold;">print</span><span style="color: Olive;">(</span><span style="color: #a11;">"原始形状:"</span><span style="color: Gray;">,</span> x.<span style="color: #05a;">shape</span><span style="color: Olive;">)</span><br />
<br />
<span style="color: #a50"># 沿 dim=1 取元素</span><br />
indices <span style="color: Gray;">=</span> torch.<span style="color: #05a;">tensor</span><span style="color: Olive;">(</span><span style="color: Olive;">[</span><span style="color: Olive;">[</span><span style="color: Maroon;">0</span><span style="color: Gray;">,</span> <span style="color: Maroon;">1</span><span style="color: Gray;">,</span> <span style="color: Maroon;">2</span><span style="color: Olive;">]</span><span style="color: Gray;">,</span><br />
<span style="color: Olive;">[</span><span style="color: Maroon;">2</span><span style="color: Gray;">,</span> <span style="color: Maroon;">0</span><span style="color: Gray;">,</span> <span style="color: Maroon;">1</span><span style="color: Olive;">]</span><span style="color: Olive;">]</span><span style="color: Olive;">)</span><br />
y <span style="color: Gray;">=</span> torch.<span style="color: #05a;">take_along_dim</span><span style="color: Olive;">(</span>x<span style="color: Gray;">,</span> indices<span style="color: Gray;">,</span> dim<span style="color: Gray;">=</span><span style="color: Maroon;">1</span><span style="color: Olive;">)</span><br />
<br />
<span style="color: Green;font-weight:bold;">print</span><span style="color: Olive;">(</span><span style="color: #a11;">"索引形状:"</span><span style="color: Gray;">,</span> indices.<span style="color: #05a;">shape</span><span style="color: Olive;">)</span><br />
<span style="color: Green;font-weight:bold;">print</span><span style="color: Olive;">(</span><span style="color: #a11;">"结果形状:"</span><span style="color: Gray;">,</span> y.<span style="color: #05a;">shape</span><span style="color: Olive;">)</span><br />
<span style="color: Green;font-weight:bold;">print</span><span style="color: Olive;">(</span><span style="color: #a11;">"n结果:"</span><span style="color: Olive;">)</span><br />
<span style="color: Green;font-weight:bold;">print</span><span style="color: Olive;">(</span>y<span style="color: Olive;">)</span><br />
</div>
</div>
<p>输出结果为:</p>
<pre>
原始形状: torch.Size([2, 3, 4])
索引形状: torch.Size([2, 3])
结果形状: torch.Size([2, 3, 4])
结果:
tensor([[[ 0, 1, 2, 3],
[ 4, 5, 6, 7],
[ 8, 9, 10, 11]],
[[16, 17, 18, 19],
[12, 13, 14, 15],
[20, 21, 22, 23]]])
Example
import torch
# Application: Selecting specific elements along the batch dimension
# For example, selecting specific key-value pairs in an attention mechanism
batch_size = 2
num_heads = 3
seq_len = 4
head_dim = 5
# Simulate the attention weights of the query
attn_weights = torch.randn(batch_size, num_heads, seq_len)
# Take the top-k indices for each head
k = 2
indices = torch.argsort(attn_weights, dim=-1, descending=True)[..., :k]
print("Index shape:", indices.shape)
# Simulate the value tensor
value = torch.randn(batch_size, num_heads, seq_len, head_dim)
# Extract the corresponding values along the seq_len dimension
selected_value = torch.take_along_dim(value, indices.unsqueeze(-1).expand(-1, -1, -1, head_dim), dim=2)
print("Value shape:", value.shape)
print("Selected Value shape:", selected_value.shape)
# Application: Selecting specific elements along the batch dimension
# For example, selecting specific key-value pairs in an attention mechanism
batch_size = 2
num_heads = 3
seq_len = 4
head_dim = 5
# Simulate the attention weights of the query
attn_weights = torch.randn(batch_size, num_heads, seq_len)
# Take the top-k indices for each head
k = 2
indices = torch.argsort(attn_weights, dim=-1, descending=True)[..., :k]
print("Index shape:", indices.shape)
# Simulate the value tensor
value = torch.randn(batch_size, num_heads, seq_len, head_dim)
# Extract the corresponding values along the seq_len dimension
selected_value = torch.take_along_dim(value, indices.unsqueeze(-1).expand(-1, -1, -1, head_dim), dim=2)
print("Value shape:", value.shape)
print("Selected Value shape:", selected_value.shape)
The output result is:
索引形状: torch.Size([2, 3, 4]) Value形状: torch.Size([2, 3, 4, 5]) 选择的Value形状: torch.Size([2, 3, 2, 5])
Note:
torch.take_along_dimIt allows indexing by dimension, which is more flexible thantorch.takebecause the latter always treats the tensor as one-dimensional.