如何仅用PyArrow对ChunkedArray进行独热编码生成PyArrow Table?
问题:仅用PyArrow将分类ChunkedArray转换为独热编码Table
现有PyArrow代码:
import pyarrow as pa ca = pa.chunked_array([['a', 'b', 'b', 'c']]) print(ca)
执行后输出:
<pyarrow.lib.ChunkedArray object at 0x7fc938bcea70> [ [ "a", "b", "b", "c" ] ]
希望得到如下格式的PyArrow Table:
pyarrow.Table _a: uint8 _b: uint8 _c: uint8 ---- _a: [[1,0,0,0]] _b: [[0,1,1,0]] _c: [[0,0,0,1]]
要求仅使用PyArrow完成(可接受依赖numpy,但优先纯PyArrow方案),避免转换为pandas。
解决方案
纯PyArrow实现方案
按以下步骤实现独热编码:
- 展平原始的嵌套ChunkedArray,方便后续处理
- 提取数组中的唯一值,作为独热编码的列名来源
- 对每个唯一值生成对应布尔数组,再转换为uint8类型的独热编码数组
- 将所有编码数组组合成目标Table
代码示例:
import pyarrow as pa # 原始ChunkedArray ca = pa.chunked_array([['a', 'b', 'b', 'c']]) # 展平嵌套数组 flat_ca = ca.flatten() # 获取排序后的唯一值 unique_vals = flat_ca.unique().sort() # 构建独热编码列 table_columns = [] for val in unique_vals: # 判断元素是否等于当前唯一值,转换为uint8 one_hot_data = flat_ca.equal(val).cast(pa.uint8()) # 包装为嵌套数组以匹配期望格式 nested_arr = pa.chunked_array([one_hot_data.to_pylist()]) # 添加列名(带_前缀)和对应数组 table_columns.append((f"_{val.as_py()}", nested_arr)) # 生成最终Table result_table = pa.Table.from_pydict(dict(table_columns)) print(result_table)
执行后输出与期望格式完全一致。
结合numpy的实现方案
如果允许使用numpy,可借助其广播机制简化编码逻辑:
import pyarrow as pa import numpy as np ca = pa.chunked_array([['a', 'b', 'b', 'c']]) # 转换为numpy数组 flat_np = ca.flatten().to_numpy() # 获取唯一值 unique_vals = np.unique(flat_np) # 生成独热编码矩阵 one_hot_matrix = (flat_np == unique_vals[:, None]).astype(np.uint8) # 转换为PyArrow数组并构建Table table_columns = [] for idx, val in enumerate(unique_vals): col_name = f"_{val}" # 包装为嵌套数组 arr = pa.chunked_array([one_hot_matrix[idx].tolist()]) table_columns.append((col_name, arr)) result_table = pa.Table.from_pydict(dict(table_columns)) print(result_table)
该方案代码更简洁,利用numpy的数组操作提升效率,但依赖numpy。
内容的提问来源于stack exchange,提问作者ignoring_gravity
相关产品推荐
相关产品推荐

