优化一个Python版 Polars函数,避免对ID进行计数或将空数据框堆叠起来
我的工作流中有一个常见模式:我有一个名为 'primary' 的数据框,可能需要对其进行行的子集化、更新数值,并可能添加新列,然后把这些子集行重新合并回主数据框。
我有一个函数,它接收一个名为 'primary' 的数据框:df1,以及一个子集数据框:df2,返回合并后的结果,方法是统计ID列的出现次数并跟踪输入数据框。可能存在数百万个唯一的ID,且df2永远不会引入新的ID。
这个函数可以从输入到输出保持完全惰性,这对我复杂的工作流很有利。但是,函数缓慢的原因是在过滤之前对堆叠后的数据框执行 pl.len().over(id),我已经通过一些性能分析确认了这一点。
# function that updates existing rows in df1 given updated rows in df2
def row_updater(df1, df2, id=pl.col("ID"), return_sorted=False, lazy=False):
# add source dataframe column
df1 = df1.lazy()
df2 = df2.lazy().with_columns(df_id = True)
# concatenate dataframes
df_result_tmp = (pl.concat([df1, df2], how="diagonal_relaxed", parallel=True).with_columns(
# count number of occurances per ID
id_count = pl.len().over(id)
)
# filter to unupdated rows from df1 and updated rows from df2
.filter((pl.col("id_count") == 1) | ((pl.col("id_count") == 2) & (pl.col("df_id")))
).drop("id_count", "df_id")
)
# sort output by ID if requested
if return_sorted:
df_result_tmp = df_result_tmp.sort(id)
# return a lazyframe if requested
if lazy:
return df_result_tmp
else:
return df_result_tmp.collect().rechunk()
我尝试使用反连接在拼接前识别唯一行,这样要快得多。只是当df1的每一行都被df2更新时,会出现Polars对将一个空帧传递给 pl.concat() 时抛出的panic异常。这可以通过在拼接前收集该帧并检查其行数来避免,但这就丢失了保持函数完全惰性的好处。使用 df.update() 的类似工作流也会失败,因为它在内部调用 pl.concat()。
# function that updates existing rows in df1 given updated rows in df2
def row_updater(df1, df2, id=pl.col("InstID"), return_sorted=False, lazy=False):
# add source dataframe column
df1 = df1.lazy()
df2 = df2.lazy().
# anti-join: keep rows in df1 whose id is NOT in df2
# df1_remaining is potentially empty
df1_remaining = df1.join(df2.select(pl.col(id)), on=id_col, how="anti")
# concatenate dataframes
# pl.concat errors when df1_remaining is empty
df_result_tmp = (pl.concat([df1_remaining, df2], how="diagonal_relaxed", parallel=True))
# sort output by ID if requested
if return_sorted:
df_result_tmp = df_result_tmp.sort(id)
# return a lazyframe if requested
if lazy:
return df_result_tmp
else:
return df_result_tmp.collect().rechunk()
解决方案
你的当前方法很慢,因为 pl.len().over(id) 会在将两个数据框拼接后强制Polars对数百万个ID进行分组和计数。一个更快的方法是使用反连接从df1中移除在df2中存在的ID,然后将剩余的行与df2拼接。
站内所有文章版权归属LeftHeroAI导航站,无授权禁止任何主体转载、抄袭、复制内容,亦不得私自架设镜像站点。一经侵权,本站将通过法律途径追责。