如何读取存储数据集的.pth文件并将其转换为CSV格式?
.pth格式数据集转CSV可行方案
完全可以实现。.pth是PyTorch生态的通用序列化文件格式,不限制存储内容类型,不管是模型权重、张量还是普通结构化数据集,都可以通过标准接口读取后转换为CSV,具体操作流程如下:
- 前置依赖安装
先安装需要的两个依赖库,执行命令:pip install torch pandas - 加载.pth文件
用PyTorch内置的load接口读取文件即可,如果文件是在GPU环境下保存、当前运行环境没有CUDA,需要加map_location参数避免设备不匹配报错:
import torch # 替换括号内的路径为你本地.pth文件的实际存储路径 pth_data = torch.load("your_dataset.pth", map_location=torch.device("cpu"))
- 确认数据结构
加载完成后先执行print(type(pth_data))查看数据类型,再打印前若干条内容确认结构,仓库附带的数据集类.pth通常是以下几类结构:- 直接存储的Pandas DataFrame对象
- 多阶PyTorch张量(Tensor)
- 键为字段名、值为列数据的字典
- 封装好的PyTorch Dataset类实例
- 嵌套列表/元组格式的二维样本数据
- 按对应结构转换导出CSV
参考以下适配不同结构的代码完成转换:
import pandas as pd import numpy as np # 情况1:加载后直接是DataFrame if isinstance(pth_data, pd.DataFrame): df = pth_data # 情况2:加载后是Tensor张量 elif isinstance(pth_data, torch.Tensor): np_array = pth_data.cpu().numpy() # 列名可根据数据集实际含义替换,以下为默认生成的特征+标签列名 col_names = [f"feat_{i}" for i in range(np_array.shape[1]-1)] + ["label"] df = pd.DataFrame(np_array, columns=col_names) # 情况3:加载后是字段对齐的字典 elif isinstance(pth_data, dict): df = pd.DataFrame(pth_data) # 情况4:加载后是PyTorch Dataset实例 elif hasattr(pth_data, "__len__") and hasattr(pth_data, "__getitem__"): sample_list = [] for idx in range(len(pth_data)): sample = pth_data[idx] # 常规分类/回归数据集的Dataset通常返回(特征, 标签)元组 if isinstance(sample, tuple) and len(sample) == 2: feat, label = sample if isinstance(feat, torch.Tensor): feat = feat.cpu().numpy().tolist() if isinstance(label, torch.Tensor): label = label.item() # 把特征和标签拼接为单行 if isinstance(feat, list): sample_list.append(feat + [label]) else: sample_list.append([feat, label]) # 同样按需替换列名 col_names = [f"feat_{i}" for i in range(len(sample_list[0])-1)] + ["label"] df = pd.DataFrame(sample_list, columns=col_names) # 导出为CSV文件,utf-8-sig编码可以避免Excel打开中文乱码 df.to_csv("converted_dataset.csv", index=False, encoding="utf-8-sig")
注意:如果.pth内存储的是图片、音频这类非结构化二进制数据,无法直接转换为规整的二维CSV表,需要先根据数据集的实际存储规则做结构化拆解后再导出,没有通用的一键转换脚本,必须先确认加载后的数据结构再做适配调整。
内容的提问来源于stack exchange,提问作者acedone
相关产品推荐
相关产品推荐

