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

解决多标签不平衡分类中的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 14:06:54