torch.searchsorted

torch.searchsorted

官方文档链接:torch.searchsorted --- PyTorch 2.3 documentation

该函数用于在已排序的序列中查找要插入的值的位置,以保持序列的顺序,

复制代码
torch.searchsorted(sorted_sequence, values, *, out_int32=False, right=False, side=None, out=None, sorter=None) → Tensor

参数如下,

  • sorted_sequence:这是一个N-D或1-D的张量,其中包含按最内部维度单调递增的序列。如果提供了sorter参数,则序列不需要按顺序排列

  • values:这是一个N-D张量或标量,包含要搜索的值

  • out_int32:这是一个可选参数,用于指示输出数据类型。如果为True,则输出数据类型为torch.int32,否则为torch.int64

  • right:这是一个可选参数,如果为False,则返回找到的第一个合适位置。如果为 True,则返回最后一个索引。如果找不到合适的索引,则对于非数值值(例如nan、inf),返回0,或者返回sorted_sequence内最内部维度的大小(超过最内部维度的最后一个索引)。如果为False,则获取每个值在sorted_sequence相应内部维度上的下限索引,如果为True,则获取上限索引。默认值为False

  • side:这是一个可选参数,"left" 对应于right为 False,"right" 对应于right为 True。如果将其设置为 "left",而right为 True,则会报错。默认值为None。

  • out:这是一个可选参数,输出张量,如果提供,则必须与 values 的大小相同

  • sorter:这是一个可选参数,如果提供,则是一个与未排序的sorted_sequence形状相匹配的张量,其中包含一个按最内部维度升序排列的索引序列

使用示例如下,

复制代码
sorted_sequence = torch.tensor([[1, 3, 5, 7, 9], [2, 4, 6, 8, 10]])
"""
tensor([[ 1,  3,  5,  7,  9],
        [ 2,  4,  6,  8, 10]])
"""

values = torch.tensor([[3, 6, 9], [3, 6, 9]])
"""
tensor([[3, 6, 9],
        [3, 6, 9]])
"""

torch.searchsorted(sorted_sequence, values)
"""
tensor([[1, 3, 4],
        [1, 2, 4]])
对于第一行 [3, 6, 9]:
数字3在第一行的sorted_sequence中的位置是索引1
数字6在第一行的sorted_sequence中的位置是索引3(6大于5而小于7,因此将6插入到索引3的位置时,能够使序列保持升序排序)
数字9在第一行的sorted_sequence中的位置是索引4
对于第二行 [3, 6, 9]:
数字3在第二行的sorted_sequence中的位置是索引1(3大于2而小于4,因此当索引为1时,不会改变序列的升序排序)
数字6在第二行的sorted_sequence中的位置是索引2
数字9在第二行的sorted_sequence中的位置是索引4(9大于8而小于10,因此当索引为4时,不会改变序列的升序排序)
"""

## 当side='right'时, 函数会返回每个值在对应行的sorted_sequence中的右侧插入位置索引
torch.searchsorted(sorted_sequence, values, side='right')
"""
tensor([[2, 3, 5],
        [1, 3, 4]])

对于第一行 [3, 6, 9]:
数字3在第一行的sorted_sequence中的右侧插入位置是索引2(数字3的右侧插入位置索引是2)
数字6在第一行的sorted_sequence中的右侧插入位置是索引3
数字9在第一行的sorted_sequence中的右侧插入位置是索引5(数字9的右侧插入位置索引是5)
对于第二行 [3, 6, 9]:
数字3在第二行的sorted_sequence中的右侧插入位置是索引1
数字6在第二行的sorted_sequence中的右侧插入位置是索引3(数字6的右侧插入位置索引是3)
数字9在第二行的sorted_sequence中的右侧插入位置是索引4
"""

sorted_sequence_1d = torch.tensor([1, 3, 5, 7, 9])
"""
tensor([1, 3, 5, 7, 9])
"""

torch.searchsorted(sorted_sequence_1d, values)
"""
tensor([[1, 3, 4],
        [1, 3, 4]])
"""
相关推荐
星核0penstarry10 分钟前
ToolGrad:把数据生成倒过来,工具调用样本通过率提到 99.8%
人工智能·测试工具·llm·函数调用·数据合成·文本调用
南京兴帝文化传媒有限公司13 分钟前
基于地图平台的本地商户信息优化:药店夜间服务标注与客户转化实操
前端·javascript·数据库·人工智能·geo 优化·geo优化避坑·ai搜索获客
tachibana219 分钟前
WebSocket 和 SSE 通信的区别及局限性
网络·人工智能·websocket·网络协议·ai·llm·agent
weixin_4462608521 分钟前
AI 驱动的 CTF 自动解题系统部署
人工智能
Thneonl24 分钟前
拆开一道 FDE 面试题,我看到三场十年前的考试
人工智能·架构
袁天25 分钟前
我从 0 基础一个月用 AI 编程做了 4 个项目,最后把踩过的坑做成了agent项目纪律系统
人工智能
葫三生26 分钟前
三生原理与《涌现:从简单规则到复杂世界》在“简单规则生成复杂系统”核心思路上存在理论呼应?
人工智能·科技·算法·机器学习·开源
南京兴帝文化传媒有限公司32 分钟前
地图SEO与AI搜索优化结合实践:宁国摄影工作室本地获客案例分析
大数据·前端·人工智能·geo 优化·geo优化避坑
sarasuki38 分钟前
上下文快爆炸了?Agent 的两种压缩手段:微压缩 vs LLM 摘要
人工智能·ai编程
林浩杨_41 分钟前
中科大 × NUS × 美团提出Pigeon:个性化图像生成
论文阅读·人工智能