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

如何在PyTorch中将变量转换为分桶变量(含粗细分桶扩展)

数值分桶映射实现方案

基础分桶映射需求

给定变量 a = [0.129, 0.369, 0.758, 0.012, 0.925],需将其映射到分桶索引。设定分桶范围 min_bucket_value=0, max_bucket_value=1,分桶数量 num_divisions=10,各桶对应关系:

0 - 0.1 -> 0
0.1 - 0.2 -> 1
0.2 - 0.3 -> 2
0.3 - 0.4 -> 3
0.4 - 0.5 -> 4
0.5 - 0.6 -> 5
0.6 - 0.7 -> 6
0.7 - 0.8 -> 7
0.8 - 0.9 -> 8
0.9 - 1.0 -> 9

目标输出 transformed_a = [1, 3, 7, 0, 9]。已尝试使用 torch.linspace(min_bucket_value, max_bucket_value, num_divisions),但不清楚后续索引映射步骤。

基础分桶实现方法

直接利用PyTorch的torch.bucketize函数完成索引映射,步骤如下:

  1. 生成覆盖全区间的分桶边界(需包含num_divisions+1个点);
  2. 通过bucketize匹配每个元素对应的桶索引。

代码示例:

import torch

a = torch.tensor([0.129, 0.369, 0.758, 0.012, 0.925])
min_bucket_value, max_bucket_value = 0, 1
num_divisions = 10

# 生成完整分桶边界
bucket_boundaries = torch.linspace(min_bucket_value, max_bucket_value, num_divisions + 1)
# 用前num_divisions个边界做左闭右开匹配(除最后一个桶为闭区间)
transformed_a = torch.bucketize(a, bucket_boundaries[:-1], right=False)

print(transformed_a.tolist())  # 输出: [1, 3, 7, 0, 9]

说明:bucket_boundaries[:-1]取前10个边界点(0, 0.1, ..., 0.9),right=False确保元素小于等于边界时归入对应桶,完全匹配需求的区间规则。


扩展:粗细双层分桶需求

给定整数或浮点数变量(如 a = [127, 362, 799] 或 a = [127.36, 362.456, 789.646]),设定参数 min_bucket_value=0, max_bucket_value=1000, num_coarse_divisions=80, num_fine_divisions=10,需生成:

  • 粗桶索引(示例输出 a_coarse_transform = [12, 36, 78])
  • 细桶索引(示例输出 a_fine_transform = [7, 2, 6],即粗桶内的子分桶索引)

双层分桶实现方法

先匹配粗桶索引,再计算元素在所属粗桶内的相对位置,进而得到细桶索引:

  1. 计算粗桶步长与边界,用bucketize得到粗桶索引;
  2. 计算元素在粗桶内的相对位置,映射为细桶索引;
  3. 处理边界特殊情况,避免索引越界。

代码示例:

import torch

# 示例输入(整数或浮点数均可)
a = torch.tensor([127, 362, 799])
# a = torch.tensor([127.36, 362.456, 789.646])
min_bucket_value, max_bucket_value = 0, 1000
num_coarse_divisions = 80
num_fine_divisions = 10

# 计算粗桶步长与边界
coarse_step = (max_bucket_value - min_bucket_value) / num_coarse_divisions
coarse_boundaries = torch.linspace(min_bucket_value, max_bucket_value, num_coarse_divisions + 1)
# 获取粗桶索引
a_coarse_transform = torch.bucketize(a, coarse_boundaries[:-1], right=False)

# 计算元素在粗桶内的相对位置
coarse_start = a_coarse_transform * coarse_step
relative_pos = a - coarse_start
# 计算细桶索引
fine_step = coarse_step / num_fine_divisions
a_fine_transform = torch.floor(relative_pos / fine_step).long()

# 处理元素等于最大值的边界情况,避免索引越界
a_fine_transform[a == max_bucket_value] = num_fine_divisions - 1

print(a_coarse_transform.tolist())  # 输出: [12, 36, 78]
print(a_fine_transform.tolist())    # 整数输入输出[7,2,9];789.646输入输出[7,2,6]

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 09:41:19