如何检测张量值是否在目标张量集合中?代码异常排查
问题:张量列表值存在性检查错误
我有一组张量列表,想要检查列表中是否存在未包含在指定张量集合unique_values中的值,但编写的代码错误地判定所有值都不在unique_values里,以下是我的实现代码及相关数据:
初始化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)
生成时间步序列并检查的代码
# 生成包含指定时间步内Data对象的序列 for i in range(start_timestep-1, start_timestep + timestep -1): np.set_printoptions(suppress=True) id_column_tensor = torch.tensor(sequence[i].id_column) list_tensor = torch.tensor([float(f"{value:.4f}") if isinstance(value, float) else value for value in id_column_tensor]) for j in unique_values: if not torch.isin(torch.tensor(j, dtype=list_tensor.dtype), list_tensor): missing_info.append(j)
unique_values内容
unique_values: {tensor(5008.), tensor(7.), tensor(5004.), tensor(0.), tensor(3.1000), tensor(7.), tensor(11.2000), tensor(12.1000), tensor(17.), tensor(5008.), tensor(18.), tensor(1.), tensor(9.2000)}
问题原因及修正方案
问题点
- 重复张量转换:
sequence[i].id_column本身已经是张量,再次用torch.tensor()包裹可能引入不必要的类型或精度问题。 - 精度不匹配:手动将浮点数格式化为4位小数再转张量,而
unique_values中的张量是原始精度(比如3.1000是PyTorch的显示格式,实际存储值可能是3.1),导致torch.isin无法匹配。 - 低效循环检查:逐个遍历
unique_values检查,效率低且容易出错。 - 集合重复元素:
unique_values是集合但包含重复元素,增加无效检查。
修正后的代码
# 先对unique_values去重并转为统一张量 unique_tensor = torch.unique(torch.tensor(list(unique_values))) for i in range(start_timestep-1, start_timestep + timestep -1): np.set_printoptions(suppress=True) # 直接使用已有的张量,无需重复转换 id_column_tensor = sequence[i].id_column # 用torch.round统一保留4位小数,避免字符串格式化的精度误差 list_tensor = torch.round(id_column_tensor * 10000) / 10000 # 批量检查缺失值,一次性获取所有不在list_tensor中的unique值 missing_mask = ~torch.isin(unique_tensor, list_tensor) missing_info.extend(unique_tensor[missing_mask].tolist())
修正说明
- 移除不必要的张量转换,直接使用原始张量。
- 通过
torch.round统一精度,确保和unique_values中的值精度一致。 - 采用批量检查方式,提升效率同时避免循环中的重复操作。
- 对
unique_values去重,减少无效检查次数。
内容的提问来源于stack exchange,提问作者kami2024
相关产品推荐
相关产品推荐

