如何将PyTorch数据集(以AG_NEWS为例)转换为Pandas DataFrame
将PyTorch的AG_NEWS数据集转换为Pandas DataFrame
直接通过以下步骤就能把AG_NEWS数据集转成Pandas DataFrame:
- 导入所需库
from torchtext.datasets import AG_NEWS import pandas as pd
- 加载数据集并转换为DataFrame
AG_NEWS加载后返回训练集和测试集两个可迭代对象,每个元素是(标签, 文本)的元组,直接传入pd.DataFrame即可:
# 加载训练集与测试集 train_data, test_data = AG_NEWS() # 转换为DataFrame,指定列名 train_df = pd.DataFrame(train_data, columns=['label', 'text']) test_df = pd.DataFrame(test_data, columns=['label', 'text'])
- 可选:优化标签可读性
AG_NEWS的原始标签是1-4的数字,对应四个新闻类别,可以把数字标签替换成类别名称:
# 定义标签与类别映射 label_map = {1: 'World', 2: 'Sports', 3: 'Business', 4: 'Sci/Tech'} # 替换标签 train_df['label'] = train_df['label'].map(label_map) test_df['label'] = test_df['label'].map(label_map)
- 可选:合并训练集与测试集
如果需要把训练和测试数据放在同一个DataFrame里,可以用pd.concat合并并标记数据拆分类型:
combined_df = pd.concat( [train_df.assign(split='train'), test_df.assign(split='test')], ignore_index=True )
内容的提问来源于stack exchange,提问作者Zenvega
相关产品推荐
相关产品推荐

