如何导出Pandas DataFrame至CSV并保留category类型
如何在CSV导出/导入时保留Pandas的Categorical类型?
CSV是纯文本格式,本身不会存储Pandas DataFrame的元数据(比如Categorical列的类别集合),直接导出再导入会丢失类型信息。以下是几种无需序列化工具的解决方案:
方法1:额外保存类别映射(推荐用于需保留完整类别集合的场景)
这种方法单独记录Categorical列的类别信息,确保导入时完全还原原类型,适合机器学习中需要固定类别编码的场景。
导出步骤
import pandas as pd import json # 示例DataFrame df = pd.DataFrame({ 'color': pd.Categorical(['red', 'blue', 'red', 'green'], categories=['red', 'blue', 'green', 'yellow']), 'size': pd.Categorical(['S', 'M', 'L', 'M'], categories=['S', 'M', 'L', 'XL']) }) # 导出核心数据到CSV df.to_csv('ml_data.csv', index=False) # 收集所有Categorical列的类别信息 cat_metadata = {col: df[col].cat.categories.tolist() for col in df.columns if df[col].dtype == 'category'} # 将类别信息保存为JSON文件 with open('cat_metadata.json', 'w') as f: json.dump(cat_metadata, f)
导入步骤
# 读取CSV数据 df_loaded = pd.read_csv('ml_data.csv') # 读取预先保存的类别信息 with open('cat_metadata.json', 'r') as f: cat_metadata = json.load(f) # 转换指定列为Categorical,还原原类别集合 for col, categories in cat_metadata.items(): df_loaded[col] = pd.Categorical(df_loaded[col], categories=categories) # 验证:可正常访问.cat属性 print(df_loaded['color'].cat.categories) # 输出: Index(['red', 'blue', 'green', 'yellow'], dtype='object')
方法2:导入时直接指定dtype(适合简单场景)
如果你的类别集合仅包含CSV中出现的所有唯一值,不需要保留未出现在数据中的类别,可以直接在read_csv时指定dtype参数:
df_loaded = pd.read_csv('ml_data.csv', dtype={'color': 'category', 'size': 'category'})
注意:这种方式会自动将CSV中该列的所有唯一值设为类别,不会保留原DataFrame中存在但未出现在当前CSV里的类别(比如示例中的yellow和XL会丢失)。
方法3:用CSV注释行存储类别信息(无需额外文件)
可以在CSV开头添加注释行来记录类别信息,这样所有数据和元数据都在一个文件里:
导出步骤
# 生成类别注释行,格式为 # cat:列名=类别1,类别2,... comment_lines = [] for col in df.columns: if df[col].dtype == 'category': category_str = ','.join(df[col].cat.categories) comment_lines.append(f'# cat:{col}={category_str}') # 先写入注释行,再追加DataFrame内容 with open('ml_data_with_comments.csv', 'w') as f: f.write('\n'.join(comment_lines) + '\n') df.to_csv(f, index=False)
导入步骤
cat_metadata = {} comment_count = 0 # 先读取注释行提取类别信息 with open('ml_data_with_comments.csv', 'r') as f: for line in f: stripped_line = line.strip() if stripped_line.startswith('# cat:'): _, content = stripped_line.split(':') col_name, categories = content.split('=') cat_metadata[col_name] = categories.split(',') comment_count += 1 else: break # 遇到数据行,停止读取注释 # 读取CSV时跳过注释行 df_loaded = pd.read_csv('ml_data_with_comments.csv', skiprows=comment_count) # 转换列为Categorical for col, categories in cat_metadata.items(): df_loaded[col] = pd.Categorical(df_loaded[col], categories=categories)
内容的提问来源于stack exchange,提问作者Nstab
相关产品推荐
相关产品推荐

