分类数据在性能与内存方面有显著优势,但使用不当也会引入陷阱。本节总结优化策略与最佳实践。
1. 内存优化原理
分类数据底层使用整数编码存储每个值,并维护一张“类别 → 编码”的映射表。当重复值较多时,内存占用远小于 object 类型。
import sys
import pandas as pd
# 构造重复率高的列
data = ['北京'] * 100000 + ['上海'] * 100000 + ['广州'] * 100000
s_obj = pd.Series(data, dtype='object')
s_cat = pd.Series(data, dtype='category')
print(f"object 内存: {s_obj.memory_usage(deep=True) / 1024:.1f} KB")
print(f"category 内存: {s_cat.memory_usage(deep=True) / 1024:.1f} KB")适用场景
- 列基数(唯一值数量)低(如小于总数 1%)
- 字符串长度较长且重复多
- 用于 groupby、value_counts 等高频操作
不适用场景
- 唯一值比例很高(如用户 ID、订单号)
- 频繁更新数据(分类新增类别需要额外处理)
2. 检查内存占用
df['城市'].astype('category').memory_usage(deep=True)
df.info(memory_usage='deep')3. 使用分类数据的最佳实践
3.1 尽早转换
在读取数据后立即将低基数列转为 category,可加速后续操作:
import pandas as pd
# 读取后直接转换
df = pd.read_csv('data.csv')
df['城市'] = df['城市'].astype('category')
# 或使用 convert_dtypes?注意 convert_dtypes 不自动转 category,需手动指定。3.2 定义复用 dtype
from pandas import CategoricalDtype
city_type = CategoricalDtype(categories=['北京', '上海', '广州'], ordered=False)
# 在多个 DataFrame 中使用
df1['城市'] = df1['城市'].astype(city_type)
df2['城市'] = df2['城市'].astype(city_type)3.3 清理未使用类别
df['城市'] = df['城市'].cat.remove_unused_categories()3.4 明确缺失值处理
# 填充时先添加类别
df['城市'] = df['城市'].cat.add_categories(['未知'])
df['城市'] = df['城市'].fillna('未知')4. 与机器学习集成
标签编码(factorize)
s = pd.Series(['北京', '上海', '北京'])
codes, uniques = pd.factorize(s)
print(codes) # [0 1 0]One-Hot 编码
df = pd.DataFrame({'城市': ['北京', '上海']})
df_encoded = pd.get_dummies(df['城市'], prefix='城市')scikit-learn 的 LabelEncoder / OneHotEncoder
from sklearn.preprocessing import LabelEncoder
le = LabelEncoder()
df['城市_编码'] = le.fit_transform(df['城市'])pandas 分类与 sklearn 互转
pd.Categorical.codes可直接作为标签编码pd.get_dummies结果可作为 OneHotEncoder 的输出
5. 常见陷阱与规避
陷阱 1:fillna 报错
s = pd.Series(['a', None]).astype('category')
# s.fillna('z') # ValueError: fill value must be in categories
s.fillna('a') # 正确:值必须在类别中陷阱 2:合并后类型退化
s1 = pd.Series(['a', 'b']).astype('category')
s2 = pd.Series(['c', 'd']).astype('category')
pd.concat([s1, s2]).dtype # object,因为类别不同解决:先统一 categories 再 concat。
s1 = s1.cat.add_categories(['c', 'd'])
s2 = s2.cat.add_categories(['a', 'b'])
pd.concat([s1, s2]).dtype # category陷阱 3:无序分类比较
s = pd.Series(['a', 'b']).astype('category')
# s > 'a' # TypeError
s.cat.as_ordered() > 'a' # 正确陷阱 4:value_counts 出现意外类别
s = pd.Series(['a', 'b']).astype(pd.CategoricalDtype(categories=['a', 'b', 'c']))
s.value_counts()
# a 1
# b 1
# c 0 ← 未出现的类别也出现陷阱 5:groupby 中出现空组
# pandas 2.x observed 默认 False
df.groupby('城市')['值'].sum() # 会包含未出现的城市(空组)
# 使用 observed=True 可只显示实际出现的组6. 性能对比示例
import pandas as pd
import numpy as np
# 构造 100 万行数据
n = 1_000_000
df = pd.DataFrame({
'类别': np.random.choice(['A', 'B', 'C'], size=n),
'数值': np.random.randn(n)
})
# 转分类
df['类别'] = df['类别'].astype('category')
# 比较内存
print(df.memory_usage(deep=True))
# 比较 groupby 速度(可实际测试)
%timeit df.groupby('类别')['数值'].mean()7. 最佳实践总结
| 实践 | 说明 |
|---|---|
| 低基数转 category | 重复率高时立即转换 |
| 复用 CategoricalDtype | 统一多列/多表结构 |
| 及时清理未用类别 | 避免内存与排序开销 |
| 注意 observed 参数 | 控制 groupby 空组输出 |
| 统一类别后再合并 | 防止类型退化 |
| 填充前确保类别存在 | 先 add_categories 再 fillna |
| 有序分类再比较 | 需要比较时先 as_ordered |
| 结合 sklearn | 直接使用 codes 或 get_dummies |
补充
分类数据并非万能:当数据频繁更新或类别数量接近样本数时,分类反而可能降低性能。评估时应结合具体数据的基数与操作模式。