Polars中如何高效实现行向点积?原生表达式尝试遇问题
Polars行向点积计算问题
问题背景
我有一个Polars DataFrame,包含values和weights两列,数据类型都是list[i64],需要对这两列执行行向点积计算。
示例DataFrame
df = pl.DataFrame({ 'values': [[0], [0, 2], [0, 2, 4], [2, 4, 0], [4, 0, 8]], 'weights': [[3], [2, 3], [1, 2, 3], [1, 2, 3], [1, 2, 3]] })
现有可行方案(性能较低)
通过struct配合map_elements实现,但Polars文档说明map_elements性能远低于原生表达式:
df.with_columns( pl.struct(['values', 'weights']) .map_elements( lambda x: np.dot(x['values'], x['weights']), return_dtype=pl.Float64 ).alias('dot') )
尝试的原生表达式方案(结果错误)
我尝试用原生表达式实现,但结果不符合预期:
df.with_columns( pl.concat_list('values', 'weights').alias('combined'), pl.concat_list('values', 'weights').list.eval(pl.element().slice(0, pl.len() // 2)).alias('values1'), pl.concat_list('values', 'weights').list.eval(pl.element().slice(pl.len() // 2, pl.len() // 2)).alias('values2'), pl.concat_list('values', 'weights').list.eval( pl.element().slice(0, pl.len() // 2).dot(pl.element().slice(pl.len() // 2, pl.len() // 2)) ).list.first().alias('dot'), pl.concat_list('values', 'weights').list.eval( pl.element().slice(0, pl.len() // 2) + pl.element().slice(pl.len() // 2, pl.len() // 2) ).alias('sum'), )
错误结果
期望dot列结果为[0, 6, 16, 10, 28],但实际输出:
shape: (5, 7) ┌───────────┬───────────┬─────────────┬───────────┬───────────┬─────┬────────────┐ │ values ┆ weights ┆ combined ┆ values1 ┆ values2 ┆ dot ┆ sum │ │ --- ┆ --- ┆ --- ┆ --- ┆ --- ┆ --- ┆ --- │ │ list[i64] ┆ list[i64] ┆ list[i64] ┆ list[i64] ┆ list[i64] ┆ i64 ┆ list[i64] │ ╞═══════════╪═══════════╪═════════════╪═══════════╪═══════════╪═════╪════════════╡ │ [0] ┆ [3] ┆ [0, 3] ┆ [0] ┆ [3] ┆ 0 ┆ [0] │ │ [0, 2] ┆ [2, 3] ┆ [0, 2, … 3] ┆ [0, 2] ┆ [2, 3] ┆ 4 ┆ [0, 4] │ │ [0, 2, 4] ┆ [1, 2, 3] ┆ [0, 2, … 3] ┆ [0, 2, 4] ┆ [1, 2, 3] ┆ 20 ┆ [0, 4, 8] │ │ [2, 4, 0] ┆ [1, 2, 3] ┆ [2, 4, … 3] ┆ [2, 4, 0] ┆ [1, 2, 3] ┆ 20 ┆ [4, 8, 0] │ │ [4, 0, 8] ┆ [1, 2, 3] ┆ [4, 0, … 3] ┆ [4, 0, 8] ┆ [1, 2, 3] ┆ 80 ┆ [8, 0, 16] │ └───────────┴───────────┴─────────────┴───────────┴───────────┴─────┴────────────┘
注意sum列结果也不符合预期,看起来是第一个切片和自身相加,而非与第二个切片相加。
问题解答
1. 错误原因分析
你在list.eval中使用pl.element()时,两次调用的pl.element()指向的是同一个列表元素上下文,导致slice操作都是基于同一个列表的切片进行计算,而非分别取values和weights的部分。简单来说,pl.element().slice(a) + pl.element().slice(b)实际上是把同一个切片的元素相加,而不是两个不同切片的对应元素相加。
2. 最优原生表达式实现
Polars提供了更直接的原生方法来处理行向列表点积,推荐三种高效方案:
方案一:pl.zip + 乘积求和(Polars >= 0.19.0)
利用pl.zip将两列列表按行配对,再计算对应元素的乘积和:
df.with_columns( pl.zip('values', 'weights') .list.eval(pl.element().first() * pl.element().second()) .list.sum() .alias('dot') )
方案二:struct列表 + 乘积求和(兼容旧版本)
将两列转为struct列表,逐元素计算乘积后求和:
df.with_columns( pl.struct(['values', 'weights']) .list.eval(pl.element().values * pl.element().weights) .list.sum() .alias('dot') )
方案三:list.zip + 索引取值求和(更简洁)
df.with_columns( pl.list.zip('values', 'weights') .list.eval(pl.element()[0] * pl.element()[1]) .list.sum() .alias('dot') )
验证结果
以上方案都能得到预期的dot列结果:[0, 6, 16, 10, 28],且性能远优于map_elements方案。
内容的提问来源于stack exchange,提问作者jackaixin
相关产品推荐
相关产品推荐

