非严格递增序列下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
相关产品推荐
相关产品推荐

