You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

非严格递增序列下torch.bucketize的替代PyTorch函数咨询

含重复节点的插值索引查找方案

关于torch.bucketize的风险

首先明确:torch.bucketize在非严格递增序列下的行为完全未定义,哪怕测试时首尾结果看起来正常,中间重复节点的场景也可能返回错误索引,绝对不能依赖这种不确定的行为来写生产代码。

替代方案

方案1:预处理节点+索引映射(适配样条首尾重复场景)

样条的重复节点大多集中在首尾,我们可以先提取唯一递增节点,用bucketize得到临时索引后,再映射回原序列的正确位置:

import torch

if __name__ == "__main__":
    boundaries = torch.tensor([0.0, 0.0, 0.0, 0.0, 0.5, 1.0, 1.0, 1.0, 1.0])
    x_test = torch.tensor([-0.1, 0.0, 0.25, 0.5, 0.75, 1.0, 1.1], dtype=torch.float64) 

    # 提取唯一递增节点
    unique_knots, _ = torch.unique(boundaries, sorted=True)
    # 构建每个唯一节点对应的原序列中最后一个出现的索引
    max_idx_map = torch.tensor([torch.where(boundaries == k)[0].max() for k in unique_knots])
    
    # 用唯一节点执行bucketize
    temp_indices = torch.bucketize(x_test, unique_knots, right=False)
    # 处理超出右边界的情况(temp_indices等于unique_knots长度时,对应原序列最后一个索引)
    final_indices = torch.where(
        temp_indices == len(unique_knots),
        len(boundaries)-1,
        max_idx_map[temp_indices]
    )

    print("原节点:    ", boundaries.tolist())
    print("测试点:    ", x_test.tolist())
    print("最终索引:  ", final_indices.tolist())

这个方法利用了样条节点重复的规律,既复用了torch.bucketize的高效实现,又保证索引完全符合插值需求。

方案2:手动实现支持重复元素的二分查找

如果不想做预处理,可以自己实现兼容重复节点的二分查找逻辑,和原torch.bucketize的行为对齐(right=False时返回第一个大于x的索引):

import torch

def bucketize_with_duplicates(x, knots, right=False):
    left = torch.zeros_like(x, dtype=torch.long)
    right_bound = torch.full_like(x, len(knots), dtype=torch.long)
    
    while (right_bound - left).any() > 0:
        mid = (left + right_bound) // 2
        # 根据right参数判断比较逻辑,处理重复元素
        if right:
            mask = x >= knots[mid]
        else:
            mask = x > knots[mid]
        left = torch.where(mask, mid + 1, left)
        right_bound = torch.where(mask, right_bound, mid)
    return left

if __name__ == "__main__":
    boundaries = torch.tensor([0.0, 0.0, 0.0, 0.0, 0.5, 1.0, 1.0, 1.0, 1.0])
    x_test = torch.tensor([-0.1, 0.0, 0.25, 0.5, 0.75, 1.0, 1.1], dtype=torch.float64) 

    indices = bucketize_with_duplicates(x_test, boundaries, right=False)
    print("原节点:    ", boundaries.tolist())
    print("测试点:    ", x_test.tolist())
    print("自定义函数索引:  ", indices.tolist())

这个实现完全可控,不管节点哪里出现重复,都能返回稳定的正确索引。

内容的提问来源于stack exchange,提问作者Mathieu

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.12 20:35:19