PyTorch torch.take_along_dim Function


Pytorch torch 参考手册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)

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)

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;">&#40;</span><span style="color: Maroon;">24</span><span style="color: Olive;">&#41;</span>.<span style="color: #05a;">reshape</span><span style="color: Olive;">&#40;</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;">&#41;</span><br />
<span style="color: Green;font-weight:bold;">print</span><span style="color: Olive;">&#40;</span><span style="color: #a11;">&quot;原始形状:&quot;</span><span style="color: Gray;">,</span> x.<span style="color: #05a;">shape</span><span style="color: Olive;">&#41;</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;">&#40;</span><span style="color: Olive;">&#91;</span><span style="color: Olive;">&#91;</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;">&#93;</span><span style="color: Gray;">,</span><br />
&nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; <span style="color: Olive;">&#91;</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;">&#93;</span><span style="color: Olive;">&#93;</span><span style="color: Olive;">&#41;</span><br />
y <span style="color: Gray;">=</span> torch.<span style="color: #05a;">take_along_dim</span><span style="color: Olive;">&#40;</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;">&#41;</span><br />
<br />
<span style="color: Green;font-weight:bold;">print</span><span style="color: Olive;">&#40;</span><span style="color: #a11;">&quot;索引形状:&quot;</span><span style="color: Gray;">,</span> indices.<span style="color: #05a;">shape</span><span style="color: Olive;">&#41;</span><br />
<span style="color: Green;font-weight:bold;">print</span><span style="color: Olive;">&#40;</span><span style="color: #a11;">&quot;结果形状:&quot;</span><span style="color: Gray;">,</span> y.<span style="color: #05a;">shape</span><span style="color: Olive;">&#41;</span><br />
<span style="color: Green;font-weight:bold;">print</span><span style="color: Olive;">&#40;</span><span style="color: #a11;">&quot;n结果:&quot;</span><span style="color: Olive;">&#41;</span><br />
<span style="color: Green;font-weight:bold;">print</span><span style="color: Olive;">&#40;</span>y<span style="color: Olive;">&#41;</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)

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.


Pytorch torch 参考手册Pytorch torch Reference Manual