如何用pandas自动转换dict为二进制DataFrame?或优化现有实现?
问题:将字典转换为one-hot格式的pandas DataFrame
我有如下字典:
d = { "anna": ["apple", "strawberry", "banana"], "bob": ["strawberry", "banana", "peach"], "chris": ["apple", "banana", "peach", "mango"] }
希望转换为如下one-hot格式的DataFrame:
apple banana mango peach strawberry anna 1 1 0 0 1 bob 0 1 0 1 1 chris 1 1 1 1 0
我已经用Python实现了该功能,但想知道pandas是否有内置方法可以自动完成转换,或者能否优化现有实现。
当前实现代码
import numpy as np import pandas as pd d = { "anna": ["apple", "strawberry", "banana"], "bob": ["strawberry", "banana", "peach"], "chris": ["apple", "banana", "peach", "mango"] } fruits = sorted(set(np.hstack(d.values()))) df = pd.DataFrame(columns=fruits) for client, client_fruits in d.items(): s = pd.Series({ fruit: fruit in client_fruits for fruit in fruits }).astype(int) df = pd.concat([df, pd.DataFrame({client: s}).T]) print(df)
解决方案:使用pandas内置方法简化实现
方法1:pd.Series.explode() + pd.crosstab()
这是最简洁的内置方法组合,步骤清晰且高效:
- 将字典转为Series,索引为用户名,值为水果列表
- 用
explode()把列表拆分为每行一个水果的结构 - 用
crosstab()生成用户与水果的交叉统计矩阵,自动得到one-hot结果
代码示例:
import pandas as pd d = { "anna": ["apple", "strawberry", "banana"], "bob": ["strawberry", "banana", "peach"], "chris": ["apple", "banana", "peach", "mango"] } # 拆分列表为多行结构 s = pd.Series(d).explode() # 生成交叉表得到one-hot矩阵 df = pd.crosstab(s.index, s.values).astype(int) # 按水果名称排序列(和目标格式一致) df = df.sort_index(axis=1) print(df)
方法2:pd.Series.explode() + pd.get_dummies()
拆分列表后用get_dummies()生成哑变量,再按用户分组求和,同样能得到目标结果:
import pandas as pd d = { "anna": ["apple", "strawberry", "banana"], "bob": ["strawberry", "banana", "peach"], "chris": ["apple", "banana", "peach", "mango"] } s = pd.Series(d).explode() df = pd.get_dummies(s).groupby(s.index).sum() df = df.sort_index(axis=1) print(df)
现有代码的优化方向
你的原代码存在**循环内多次pd.concat()**的问题,这在数据量大时会产生大量临时对象,性能极低。可以优化为:
- 预先创建全0的DataFrame,直接指定索引为用户名、列为水果
- 遍历每个用户,批量给对应水果列赋值1
优化后的代码:
import numpy as np import pandas as pd d = { "anna": ["apple", "strawberry", "banana"], "bob": ["strawberry", "banana", "peach"], "chris": ["apple", "banana", "peach", "mango"] } # 获取所有水果并排序 fruits = sorted(set(np.hstack(d.values()))) # 初始化全0的DataFrame df = pd.DataFrame(0, index=d.keys(), columns=fruits) # 批量赋值,避免循环concat for user, fruit_list in d.items(): df.loc[user, fruit_list] = 1 print(df)
这个版本彻底避免了低效的concat操作,性能比原代码提升显著,尤其是用户数量较多时。
内容的提问来源于stack exchange,提问作者Lewan
相关产品推荐
相关产品推荐

