如何在cuDF DataFrame中对分组聚合后的列表列进行填充?
基于cuDF/cuPy原生实现分组列表左填充的方案
问题概述
需要对cuDF DataFrame分组聚合后的列表执行左填充,将所有列表统一到指定最大长度(示例为3),填充值为-1。现有转pandas+np.pad的方法效率低,直接调用cuDF的apply会触发NumbaNotImplementedError。
需求示例
概念性目标代码:
df = cudf.DataFrame({"g": [1, 1, 1, 2, 2, 3], "a": [1, 2, 3, 1, 3, 1]}) df.groupby("g")["a"].collect().list.pad(max_length=3, pad_left=True, drop="last", padding_value=-1)
期望输出:
g 1 [1, 2, 3] 2 [-1, 1, 3] 3 [-1, -1, 1]
当前问题
- 转pandas实现繁琐且慢:
cudf.from_pandas( df.groupby("g")["a"] .collect() .to_pandas() .apply(lambda x: np.pad(x, (max(3 - len(x), 0), 0), constant_values=(-1,))) ) - 直接调用
apply报错:
错误信息:df.groupby("g")["a"].collect().apply( lambda x: np.pad(x, (max(3 - len(x), 0), 0), constant_values=(-1,)) )NumbaNotImplementedError: list
原生解决方案
方法1:cuDF列表原生操作拼接
利用cuDF的列表拼接和截取能力,避免跨设备传输:
import cudf import cupy as cp df = cudf.DataFrame({"g": [1, 1, 1, 2, 2, 3], "a": [1, 2, 3, 1, 3, 1]}) max_len = 3 pad_val = -1 # 分组得到列表 grouped_series = df.groupby("g")["a"].collect() # 计算每个分组需要填充的元素数 pad_counts = cp.maximum(max_len - grouped_series.list.len(), 0) # 生成对应长度的填充列表 pad_lists = pad_counts.apply(lambda x: [pad_val] * x) # 拼接填充列表与原列表,截取前max_len个元素(左填充) result = pad_lists.list.concat(grouped_series).list.take(slice(0, max_len)) print(result)
方法2:cuPy向量化填充(适合大数据量)
通过展开列表为二维数组,用cuPy批量填充后重新打包:
import cudf import cupy as cp df = cudf.DataFrame({"g": [1, 1, 1, 2, 2, 3], "a": [1, 2, 3, 1, 3, 1]}) max_len = 3 pad_val = -1 grouped_series = df.groupby("g")["a"].collect() lengths = grouped_series.list.len().values offsets = cp.cumsum(cp.array([0] + lengths.tolist()))[:-1] # 初始化全填充值的二维数组 full_arr = cp.full((len(grouped_series), max_len), pad_val) # 将原数据填充到数组的右侧位置(对应左填充) flat_data = grouped_series.list.flatten().to_cupy() for idx in range(len(grouped_series)): start_col = max_len - lengths[idx] full_arr[idx, start_col:] = flat_data[offsets[idx]:offsets[idx]+lengths[idx]] # 转换回cuDF列表Series result = cudf.Series(cp.split(full_arr, len(grouped_series)), index=grouped_series.index) print(result)
输出验证
两种方法均会输出符合预期的结果:
g 1 [1, 2, 3] 2 [-1, 1, 3] 3 [-1, -1, 1] dtype: list
内容的提问来源于stack exchange,提问作者bilzard
相关产品推荐
相关产品推荐

