如何在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函数完成索引映射,步骤如下:
- 生成覆盖全区间的分桶边界(需包含
num_divisions+1个点); - 通过
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],即粗桶内的子分桶索引)
双层分桶实现方法
先匹配粗桶索引,再计算元素在所属粗桶内的相对位置,进而得到细桶索引:
- 计算粗桶步长与边界,用
bucketize得到粗桶索引; - 计算元素在粗桶内的相对位置,映射为细桶索引;
- 处理边界特殊情况,避免索引越界。
代码示例:
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
相关产品推荐
相关产品推荐

