如何在Polars DataFrame中高效存储矩阵并并行执行transpose(v)*M*v运算?
用Polars高效计算按日期分组的
v^T * M * v标量结果 需求说明
存在大量日期维度的数据,每个日期对应一个长度为n的向量v和n×n的方阵M,其中v、M的取值以及n的大小会随日期变化。需要为每个日期计算矩阵运算transpose(v) * M * v的标量结果,传统循环方法效率极低,希望借助Polars DataFrame的group_by("date").agg(...)实现并行高效计算。
示例数据
- 日期
2020-01-01:v = [1, 2],M = [[1, 2], [3, 4]],向量和矩阵的索引为["a", "b"] - 日期
2020-02-01:v = [1, 2, 3],M = [[1, 2, 3], [3, 4, 5], [4, 5, 6]],索引为["a", "b", "c"] - 日期
2020-03-01:v = [1, 3, 5],M = [[1, 5, 9], [2, 4, 6], [3, 6, 9]],索引为["b", "d", "e"]
预期运算结果
2020-01-01:272020-02-01:1542020-03-01:523
Polars实现方案
步骤1:构造结构化数据
首先将每个日期的向量和矩阵数据整理为Polars可处理的长格式:
import polars as pl import numpy as np # 构造示例数据集 data = [ {"date": "2020-01-01", "idx": "a", "v": 1, "M_a": 1, "M_b": 2}, {"date": "2020-01-01", "idx": "b", "v": 2, "M_a": 3, "M_b": 4}, {"date": "2020-02-01", "idx": "a", "v": 1, "M_a": 1, "M_b": 2, "M_c": 3}, {"date": "2020-02-01", "idx": "b", "v": 2, "M_a": 3, "M_b": 4, "M_c": 5}, {"date": "2020-02-01", "idx": "c", "v": 3, "M_a": 4, "M_b": 5, "M_c": 6}, {"date": "2020-03-01", "idx": "b", "v": 1, "M_b": 1, "M_d": 5, "M_e": 9}, {"date": "2020-03-01", "idx": "d", "v": 3, "M_b": 2, "M_d": 4, "M_e": 6}, {"date": "2020-03-01", "idx": "e", "v": 5, "M_b": 3, "M_d": 6, "M_e": 9}, ] df = pl.DataFrame(data)
步骤2:分组计算核心运算
提供两种实现方式,均可利用Polars的并行分组能力提升效率:
方法1:结合Numpy向量化计算(推荐,效率更高)
利用矩阵运算的原生支持,通过group_by.apply实现分组计算:
def compute_vTm_v(group_df: pl.DataFrame) -> pl.DataFrame: # 提取当前分组的向量v和矩阵M v = group_df["v"].to_numpy() m_cols = [col for col in group_df.columns if col.startswith("M_")] M = group_df[m_cols].to_numpy() # 计算v^T * M * v result = v.T @ M @ v return pl.DataFrame({"date": [group_df["date"][0]], "result": [result]}) # 分组应用并合并结果 result_df = df.group_by("date").apply(compute_vTm_v) print(result_df)
方法2:纯Polars表达式计算(无需依赖Numpy)
通过Polars的列表和映射元素功能实现运算:
# 获取所有唯一索引,用于匹配M列 all_idx = df["idx"].unique().to_list() result_df = df.group_by("date").agg( v_list=pl.col("v").list(), # 收集每行的M值为字典,键为索引 M_rows=pl.struct([pl.col(f"M_{idx_val}").alias(idx_val) for idx_val in all_idx]).list() ).with_columns( # 展开计算sum(v_i * M_ij * v_j) result=pl.map_elements( lambda row: sum( row["v_list"][i] * row["M_rows"][i][row["v_list"].index()[j]] * row["v_list"][j] for i in range(len(row["v_list"])) for j in range(len(row["v_list"])) ), return_dtype=pl.Float64 ) ).select("date", "result") print(result_df)
最终输出结果
两种方法都会得到如下结果:
shape: (3, 2) date result --- --- str f64 2020-01-01 27.0 2020-02-01 154.0 2020-03-01 523.0
内容的提问来源于stack exchange,提问作者lebesgue
相关产品推荐
相关产品推荐

