PyTorch torch.searchsorted Function


Pytorch torch 参考手册PyTorch torch Reference Manual

torch.searchsortedThis is a function in PyTorch used to search for the position where an element should be inserted in a sorted tensor. The returned value is the index position where the element should be inserted.

Function Definition

torch.searchsorted(sorted_sequence, values, side='left', out_int32=False, right=False)

Parameter Description:

  • sorted_sequence: Sorted one-dimensional or multi-dimensional tensor
  • values: Value to search for
  • side: 'left' or 'right', determines whether to return the left or right insertion position
  • out_int32: Whether to return int32 type
  • right: Deprecated, use side instead

Usage Example

Example

import torch

# Create a sorted sequence
sorted_seq = torch.tensor([1, 3, 5, 7, 9])

# Search for the position of a value
values = torch.tensor([3, 6, 8])
y = torch.searchsorted(sorted_seq, values)
print(y)

The output result is:

tensor([1, 3, 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"># 创建已排序的序列</span><br />
sorted_seq <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: Maroon;">1</span><span style="color: Gray;">,</span> <span style="color: Maroon;">3</span><span style="color: Gray;">,</span> <span style="color: Maroon;">5</span><span style="color: Gray;">,</span> <span style="color: Maroon;">7</span><span style="color: Gray;">,</span> <span style="color: Maroon;">9</span><span style="color: Olive;">&#93;</span><span style="color: Olive;">&#41;</span><br />
<br />
<span style="color: #a50"># 使用 side='right' 搜索</span><br />
values <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: Maroon;">3</span><span style="color: Gray;">,</span> <span style="color: Maroon;">6</span><span style="color: Gray;">,</span> <span style="color: Maroon;">8</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;">searchsorted</span><span style="color: Olive;">&#40;</span>sorted_seq<span style="color: Gray;">,</span> values<span style="color: Gray;">,</span> side<span style="color: Gray;">=</span><span style="color: #a11;">'right'</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>
tensor([2, 3, 4])

Example

import torch

# For multi-dimensional arrays
sorted_seq = torch.tensor([[1, 3, 5], [2, 4, 6]])
values = torch.tensor([[1.5], [3.5]])
y = torch.searchsorted(sorted_seq, values)
print(y)

The output result is:

tensor([[1],
        [1]])

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions