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
相关产品推荐
相关产品推荐

