如何避免PyTorch Tensor改变输入格式并实现浮点数ID去重?
浮点类型ID存入PyTorch Data对象后去重失败的解决方法
问题背景
我有一组作为ID的浮点数值,原始格式是11.0、11.2、16.2这类,存入PyTorch的Data对象时用了以下代码:
data_with_id = Data(x=x_without_id, edge_index=edge_index.t().contiguous(), edge_weight=edge_feats, y=ground_truth_labels) data_with_id.id_column = torch.tensor(id_column, dtype=torch.float)
转换后发现tensor显示的格式变了,比如11.0变成tensor(11.),11.2变成tensor(11.2000),而且用集合去重时仍有重复项(比如5001.和17.都出现了两次):
unique_values: {tensor(11.), tensor(11.2000), tensor(16.2000),
tensor(6.1000), tensor(5006.), tensor(17.), tensor(13.2000),
tensor(19.1000), tensor(2.), tensor(9.), tensor(5010.), tensor(6.),
tensor(14.), tensor(14.1000), tensor(5001.), tensor(5.1000),
tensor(6.2000), tensor(5001.), tensor(4.2000), tensor(11.1000),
tensor(12.), tensor(5007.), tensor(17.1000), tensor(17.),
tensor(4.1000), tensor(8.2000),tensor(19.1000)}
我尝试了以下代码去重,但没用:
missing_info = [] unique_values = set() for i in range(start_timestep-1, start_timestep+ timestep-1): if sequence[i].x.size(0) > 0: np.set_printoptions(suppress=True) id_column_tensor = torch.tensor(sequence[i].id_column) id_column_list = [float(f"{value:.4f}") if isinstance(value, float) else value for value in id_column_tensor] unique_values.update(id_column_list) else: raise ValueError(f"The tensor at sequence[{i}].x is empty.") print('unique_values:',unique_values)
问题原因
- 浮点精度误差:浮点数在计算机中存储时可能存在微小偏差,比如5001.0可能被存成5001.0000000001,肉眼看不出来但会被集合判定为不同值。
- 数值提取错误:遍历tensor时,每个元素是tensor对象而非原生float,判断
isinstance(value, float)永远为False,导致部分数值没被正确转换就存入集合,集合无法识别tensor对象的重复。
解决方法
修正后的去重代码
missing_info = [] unique_values = set() # 根据原始ID的小数位数设置精度,这里原始ID最多1位小数,设为1即可 precision = 1 for i in range(start_timestep-1, start_timestep + timestep - 1): if sequence[i].x.size(0) > 0: # 直接使用已有的id_column tensor,无需重复转换 id_tensor = sequence[i].id_column # 统一精度,消除浮点误差 rounded_ids = torch.round(id_tensor * (10 ** precision)) / (10 ** precision) # 将tensor元素转换为原生float id_list = [item.item() for item in rounded_ids] unique_values.update(id_list) else: raise ValueError(f"sequence[{i}].x 是空张量") print('unique_values:', unique_values)
关键改动说明
- 统一浮点精度:用
torch.round对ID进行四舍五入,保留和原始数据一致的小数位数,消除存储带来的微小误差。 - 正确提取数值:用
.item()将tensor中的每个元素转换为Python原生float,确保集合能正确识别重复值。 - 避免重复转换:
sequence[i].id_column本身已经是tensor,无需再用torch.tensor()重复转换。
备选方案(用numpy处理)
如果习惯用numpy,也可以这样写:
missing_info = [] unique_values = set() precision = 1 for i in range(start_timestep-1, start_timestep + timestep - 1): if sequence[i].x.size(0) > 0: id_np = sequence[i].id_column.numpy() rounded_ids = np.round(id_np, decimals=precision) id_list = rounded_ids.tolist() unique_values.update(id_list) else: raise ValueError(f"sequence[{i}].x 是空张量") print('unique_values:', unique_values)
内容的提问来源于stack exchange,提问作者kami2024
相关产品推荐
相关产品推荐

