如何在Polars中获取每行最大值所在的列名?
在Polars中获取每行最大值对应的列名
在Pandas中,我们可以通过idxmax(axis=1)直接获取每行最大值所在的列名,但Polars没有提供完全对应的方法。以下是几种可行的实现方式:
方法1:利用arg_max + horizontal参数
Polars的arg_max函数支持horizontal参数,开启后会跨行计算最大值的索引,再将索引映射为列名即可:
import polars as pl df = pl.DataFrame({'a': [1, 2, 3, 4, 5], 'b': [5, 4, 3, 2, 1]}) df = df.with_columns( Largest=pl.arg_max(pl.all(), horizontal=True).map_elements(lambda idx: df.columns[idx]) ) print(df)
执行结果:
shape: (5, 3) ┌─────┬─────┬─────────┐ │ a ┆ b ┆ Largest │ │ --- ┆ --- ┆ --- │ │ i64 ┆ i64 ┆ str │ ╞═════╪═════╪═════════╡ │ 1 ┆ 5 ┆ b │ │ 2 ┆ 4 ┆ b │ │ 3 ┆ 3 ┆ a │ │ 4 ┆ 2 ┆ a │ │ 5 ┆ 1 ┆ a │ └─────┴─────┴─────────┘
方法2:max_horizontal + 分支判断
如果数据列数较少,可以先获取每行最大值,再通过when/then匹配对应的列名:
import polars as pl df = pl.DataFrame({'a': [1, 2, 3, 4, 5], 'b': [5, 4, 3, 2, 1]}) row_max = pl.max_horizontal(pl.all()) df = df.with_columns( Largest=pl.when(pl.col('a') == row_max).then('a').when(pl.col('b') == row_max).then('b') ) print(df)
这种方式直观但扩展性差,列数多的时候需要写大量分支。
方法3:宽表转长表 + 分组聚合
通过melt将宽表转为长表,按行分组后筛选出最大值对应的列名,最后合并回原表:
import polars as pl df = pl.DataFrame({'a': [1, 2, 3, 4, 5], 'b': [5, 4, 3, 2, 1]}) # 添加行索引用于分组关联 result = ( df.with_row_index('idx') .melt(id_vars='idx', value_name='val', variable_name='col') .group_by('idx') .agg(pl.col('col').filter(pl.col('val') == pl.col('val').max()).first()) .join(df.with_row_index('idx'), on='idx') .drop('idx') .rename({'col': 'Largest'}) ) print(result)
这个方法适合任意列数的场景,通用性更强。
内容的提问来源于stack exchange,提问作者asongtoruin
相关产品推荐
相关产品推荐

