You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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:27
  • 2020-02-01:154
  • 2020-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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.14 09:27:08