解决多标签不平衡分类中的IndexError张量维度报错问题
解决多标签分类任务中的IndexError问题
错误原因
你处理的是10标签不平衡多标签分类任务,training_loader.dataset.labels是二维结构(每个样本对应10个标签的0/1标记,形状为[样本数, 10]),但你定义的sample_probs是一维Tensor(形状[10])。当执行sample_probs[label]时,label是长度为10的数组,一维Tensor无法处理这种多维索引,因此抛出IndexError: too many indices for tensor of dimension 1。
修正方案
1. 调整多标签场景下的样本筛选逻辑
根据多标签特性,重新编写样本删除的判断逻辑(逻辑可根据你的需求调整):
# 假设labels是二维数组(numpy或tensor),shape为[N, 10],0表示无该标签,1表示有 idx_to_del = [] for i, label_vec in enumerate(training_loader.dataset.labels): # 示例逻辑:如果样本中所有存在的标签,对应的随机数都超过sample_probs,则删除该样本 should_delete = True for label_idx in range(num_classes): # 仅针对样本存在的标签进行判断 if label_vec[label_idx] == 1: if random.random() <= sample_probs[label_idx]: should_delete = False break if should_delete: idx_to_del.append(i)
2. 修复数据集加载错误
原代码中imbalanced_train_loader传入的是原始train_dataset,应改为处理后的imbalanced_train_dataset,否则之前的数据集处理无效:
imbalanced_train_loader = torch.utils.data.DataLoader( imbalanced_train_dataset, batch_size=8, shuffle=True, **kwargs )
3. 可选:统一数据类型
如果labels是Tensor类型,建议先转为numpy数组处理,避免索引冲突:
labels_np = training_loader.dataset.labels.numpy() idx_to_del = [] for i, label_vec in enumerate(labels_np): # 同上判断逻辑 ...
内容的提问来源于stack exchange,提问作者bbb
相关产品推荐
相关产品推荐

