PyTorch torch.searchsorted Function
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 tensorvalues: Value to search forside: 'left' or 'right', determines whether to return the left or right insertion positionout_int32: Whether to return int32 typeright: 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)
# 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;">(</span><span style="color: Olive;">[</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;">]</span><span style="color: Olive;">)</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;">(</span><span style="color: Olive;">[</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;">]</span><span style="color: Olive;">)</span><br /> y <span style="color: Gray;">=</span> torch.<span style="color: #05a;">searchsorted</span><span style="color: Olive;">(</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;">)</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> 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)
# 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]])
Other Extensions