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

Dataset.from_dict()加载多热编码标签维度异常求助

解决Dataset.from_dict加载多热编码标签时维度反转的问题

问题原因

当你用Dataset.from_dict加载二维列表形式的多热编码标签时,Hugging Face Dataset会默认将其识别为序列型特征(每个样本包含3个独立的float元素),而非每个样本是3维的标签向量。PyTorch DataLoader在批量堆叠时,会按序列的位置维度聚合,最终导致原本期望的(5,3)张量变成了3个长度为5的张量(形状等价于(3,5))。

修复方案

显式定义特征类型,告诉Dataset你的标签是每个样本对应一个3维向量的二维数组,使用Array2D类型来指定:

from datasets import Dataset, Features, Array2D
from torch.utils.data import DataLoader

texts = ['a', 'b', 'c', 'd', 'e']
multihot_labels= [[0.0, 0.0, 1.0], [0.0, 1.0, 1.0], [0.0, 1.0, 0.0], [0.0, 1.0, 0.0], [0.0, 0.0, 1.0]]

# 自定义特征结构,指定label为每个样本3维的数组
features = Features({
    'text': 'string',
    'label': Array2D(shape=(3,), dtype='float32')
})

dataset = Dataset.from_dict({'text': texts, 'label': multihot_labels}, features=features)
loader = DataLoader(dataset, batch_size=5)

for batch in loader:
    print(batch['text'])
    print(batch['label'])
    print(batch['label'].shape)  # 输出 torch.Size([5, 3]),符合预期
    break

补充说明

  • Array2D(shape=(3,), dtype='float32')表示每个样本的标签是一个长度为3的一维数组,批量后会自动堆叠成(batch_size, 3)的张量。
  • 如果你不想自定义特征,也可以通过自定义collate_fn来调整维度,但显式指定特征类型是更符合Dataset设计规范的做法:
    def custom_collate(batch):
        texts = [item['text'] for item in batch]
        labels = torch.tensor([item['label'] for item in batch], dtype=torch.float32)
        return {'text': texts, 'label': labels}
    
    loader = DataLoader(dataset, batch_size=5, collate_fn=custom_collate)
    

内容的提问来源于stack exchange,提问作者Sandy

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.07 23:07:12