如何在Polars中使用pl.all()遍历所有列并适配自定义水平填充函数?
如何在Polars中使用pl.all()遍历所有列并适配自定义水平填充函数?
我完全懂你碰到的这个困扰——pl.all()返回的是单个表达式对象,没法直接像迭代器那样反转或者遍历,这直接导致你没法把它作为默认参数用到原来的fill_horizontal函数里。不过咱们可以调整函数的参数处理逻辑,再结合Polars的cum_reduce方法,既能解决pl.all()的适配问题,还能让函数实现更简洁高效。
下面是调整后的函数实现,已经兼容了pl.all()作为默认参数的场景:
from typing import Iterable from polars._typing import IntoExpr import polars as pl def fill_horizontal( exprs: Iterable[IntoExpr] | None = None, *, forward: bool = True, ncols: int = 1000) -> pl.Expr: """Generate a horizontal forward/backward fill expression.""" if exprs is None: # 正向填充直接用pl.all()即可,cum_reduce能处理它 # 反向填充需要指定列的范围,ncols设为不小于实际列数的值就行 cols = pl.all() if forward else pl.nth(range(ncols, -1, -1)) else: # 传入自定义列时,根据填充方向决定是否反转 cols = exprs if forward else reversed(exprs) # 用cum_reduce做累积合并,最后拆分成单独列 return pl.cum_reduce(lambda s1, s2: pl.coalesce(s2, s1), cols).struct.unnest()
关键逻辑说明
- 默认参数处理:当
exprs为None时,正向填充直接用pl.all(),cum_reduce方法可以直接处理这个表达式,自动遍历所有列;反向填充时用pl.nth(range(ncols, -1, -1))来反向选取列,你只需要保证ncols的值不小于数据框的实际列数就好。 - 累积合并:
pl.cum_reduce会按列的顺序依次执行lambda里的coalesce操作,正好对应水平填充的逻辑——每一列都和前面的结果合并,拿到第一个非空值。 - 结果拆分:
cum_reduce返回的是结构体,用struct.unnest()可以把它拆成单独的列,直接和原数据框合并。
实际使用示例
用你提供的测试数据来验证一下效果:
df = pl.DataFrame({ "col1": [1, None, 2], "col2": [1, 2, None], "col3": [None, None, 3]}) print(df) # shape: (3, 3) # ┌──────┬──────┬──────┐ # │ col1 ┆ col2 ┆ col3 │ # │ --- ┆ --- ┆ --- │ # │ i64 ┆ i64 ┆ i64 │ # ╞══════╪══════╪══════╡ # │ 1 ┆ 1 ┆ null │ # │ null ┆ 2 ┆ null │ # │ 2 ┆ null ┆ 3 │ # └──────┴──────┴──────┘ print('正向水平填充(默认用所有列)') print(df.with_columns(fill_horizontal())) # shape: (3, 3) # ┌──────┬──────┬──────┐ # │ col1 ┆ col2 ┆ col3 │ # │ --- ┆ --- ┆ --- │ # │ i64 ┆ i64 ┆ i64 │ # ╞══════╪══════╪══════╡ # │ 1 ┆ 1 ┆ 1 │ # │ null ┆ 2 ┆ 2 │ # │ 2 ┆ 2 ┆ 3 │ # └──────┴──────┴──────┘ print('反向水平填充(默认用所有列,注意ncols要足够大)') print(df.with_columns(fill_horizontal(forward=False, ncols=3))) # shape: (3, 3) # ┌──────┬──────┬──────┐ # │ col1 ┆ col2 ┆ col3 │ # │ --- ┆ --- ┆ --- │ # │ i64 ┆ i64 ┆ i64 │ # ╞══════╪══════╪══════╡ # │ 1 ┆ 1 ┆ null │ # │ 2 ┆ 2 ┆ null │ # │ 2 ┆ 3 ┆ 3 │ # └──────┴──────┴──────┘
这样调整后,你既可以传入自定义的列列表,也可以直接使用默认的pl.all()来处理所有列,完美解决了之前的类型错误问题。
备注:内容来源于stack exchange,提问作者Olibarer
相关产品推荐
相关产品推荐

