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

如何更快找到多LazyFrame列函数最小化结果对应的x值?

优化按周期的多列函数最小化计算速度

我有一个包含多列小时数据的Polars LazyFrame,数据按周期划分。需要为每个周期找到x值,使一个涉及多列数学运算的函数结果最小化。目前用scipy.optimize.minimize实现,但运行速度极慢,求更快的实现方案。


原实现代码

import scipy
import polars as pl
from datetime import datetime

hourly_data = pl.DataFrame({
    'period': [0,0,0,0,0,0,1,1,1,1,1,1,2,2,2,2,2,2,3,3,3,3,3,3],
    'price': [4,4,4,4,4,4,5,5,5,5,5,5,6,6,6,6,6,6,7,7,7,7,7,7],
    'quantity': [7,7,7,7,7,7,6,6,6,6,6,6,5,5,5,5,5,5,4,4,4,4,4,4],
    'estimated_price': [5,5,5,5,5,5,6,6,6,6,6,6,7,7,7,7,7,7,8,8,8,8,8,8],
    'estimated_quantity': [6,6,6,6,6,6,5,5,5,5,5,5,4,4,4,4,4,4,3,3,3,3,3,3],
    'key_product': [0.9,0.8,0.7,0.8,0.9,0.8,0.7,0.8,0.9,0.8,0.7,0.8,0.9,0.8,0.7,0.8,0.9,0.8,0.7,0.8,0.9,0.8,0.7,0.8],
    'initial_guess': [10,10,10,10,10,10,20,20,20,20,20,20,30,30,30,30,30,30,40,40,40,40,40,40]
}).lazy()
hourly_data = hourly_data.with_columns(pl.datetime_range(datetime(2024,1,1), datetime(2024,1,1,23), '1h').alias('hour'))
hourly_data = hourly_data.with_columns(pl.col('hour').min().over('period').alias('period_start'))

def minimization_target(x, period_start):
    return hourly_data.filter(pl.col('period_start') == period_start).select(
        (((pl.col('price').median() * pl.col('quantity').median() - (pl.col('estimated_quantity') * (pl.col('estimated_price') + x)).sum()) 
        / (pl.col('key_product') * (pl.col('estimated_price') + x)).sum()).abs() - 1).abs()
    ).collect().item()

results = hourly_data.group_by('period_start', maintain_order=True).map_groups(
    lambda group: pl.DataFrame({'x_values': scipy.optimize.minimize(
        minimization_target, 
        group.get_column('initial_guess').median(), 
        args=group.get_column('period_start').median()
    ).x}), 
    schema=None
)

输入数据

shape: (24, 9)
┌────────┬───────┬──────────┬─────────────┬───┬─────────────┬────────────┬────────────┬────────────┐
│ period ┆ price ┆ quantity ┆ estimated_p ┆ … ┆ key_product ┆ initial_gu ┆ hour       ┆ period_sta │
│ ---    ┆ ---   ┆ ---      ┆ rice        ┆   ┆ ---         ┆ ess        ┆ ---        ┆ rt         │
│ i64    ┆ i64   ┆ i64      ┆ ---         ┆   ┆ f64         ┆ ---        ┆ datetime[μ ┆ ---        │
│        ┆       ┆          ┆ i64         ┆   ┆             ┆ i64        ┆ s]         ┆ datetime[μ │
│        ┆       ┆          ┆             ┆   ┆             ┆            ┆            ┆ s]         │
╞════════╪═══════╪══════════╪═════════════╪═══╪═════════════╪════════════╪════════════╪════════════╡
│ 0      ┆ 4     ┆ 7        ┆ 5           ┆ … ┆ 0.9         ┆ 10         ┆ 2024-01-01 ┆ 2024-01-01 │
│        ┆       ┆          ┆             ┆   ┆             ┆            ┆ 00:00:00   ┆ 00:00:00   │
│ 0      ┆ 4     ┆ 7        ┆ 5           ┆ … ┆ 0.8         ┆ 10         ┆ 2024-01-01 ┆ 2024-01-01 │
│        ┆       ┆          ┆             ┆   ┆             ┆            ┆ 01:00:00   ┆ 00:00:00   │
│ 0      ┆ 4     ┆ 7        ┆ 5           ┆ … ┆ 0.7         ┆ 10         ┆ 2024-01-01 ┆ 2024-01-01 │
│        ┆       ┆          ┆             ┆   ┆             ┆            ┆ 02:00:00   ┆ 00:00:00   │
│ 0      ┆ 4     ┆ 7        ┆ 5           ┆ … ┆ 0.8         ┆ 10         ┆ 2024-01-01 ┆ 2024-01-01 │
│        ┆       ┆          ┆             ┆   ┆             ┆            ┆ 03:00:00   ┆ 00:00:00   │
│ 0      ┆ 4     ┆ 7        ┆ 5           ┆ … ┆ 0.9         ┆ 10         ┆ 2024-01-01 ┆ 2024-01-01 │
│        ┆       ┆          ┆             ┆   ┆             ┆            ┆ 04:00:00   ┆ 00:00:00   │
│ …      ┆ …     ┆ …        ┆ …           ┆ … ┆ …           ┆ …          ┆ …          ┆ …          │
│ 3      ┆ 7     ┆ 4        ┆ 8           ┆ … ┆ 0.8         ┆ 40         ┆ 2024-01-01 ┆ 2024-01-01 │
│        ┆       ┆          ┆             ┆   ┆             ┆            ┆ 19:00:00   ┆ 18:00:00   │
│ 3      ┆ 7     ┆ 4        ┆ 8           ┆ … ┆ 0.9         ┆ 40         ┆ 2024-01-01 ┆ 2024-01-01 │
│        ┆       ┆          ┆             ┆   ┆             ┆            ┆ 20:00:00   ┆ 18:00:00   │
│ 3      ┆ 7     ┆ 4        ┆ 8           ┆ … ┆ 0.8         ┆ 40         ┆ 2024-01-01 ┆ 2024-01-01 │
│        ┆       ┆          ┆             ┆   ┆             ┆            ┆ 21:00:00   ┆ 18:00:00   │
│ 3      ┆ 7     ┆ 4        ┆ 8           ┆ … ┆ 0.7         ┆ 40         ┆ 2024-01-01 ┆ 2024-01-01 │
│        ┆       ┆          ┆             ┆   ┆             ┆            ┆ 22:00:00   ┆ 18:00:00   │
│ 3      ┆ 7     ┆ 4        ┆ 8           ┆ … ┆ 0.8         ┆ 40         ┆ 2024-01-01 ┆ 2024-01-01 │
│        ┆       ┆          ┆             ┆   ┆             ┆            ┆ 23:00:00   ┆ 18:00:00   │
└────────┴───────┴──────────┴─────────────┴───┴─────────────┴────────────┴────────────┴────────────┘

期望输出

shape: (4, 1)
┌────────────┐
│ x_values   │
│ ---        │
│ f64        │
╞════════════╡
│ -16.006287 │
│ 10.331055  │
│ 25.420471  │
│ 37.352234  │
└────────────┘

优化方案

原代码的核心问题是每次优化迭代都要重复查询、过滤Polars数据并执行collect,带来巨大的IO和计算开销。优化思路是先预计算所有周期的固定聚合值,将目标函数转化为纯数值运算。

步骤1:预计算周期级聚合参数

一次性计算每个周期所需的固定值,避免重复查询:

import scipy.optimize
import polars as pl
from datetime import datetime

# 初始化数据
hourly_data = pl.DataFrame({
    'period': [0,0,0,0,0,0,1,1,1,1,1,1,2,2,2,2,2,2,3,3,3,3,3,3],
    'price': [4,4,4,4,4,4,5,5,5,5,5,5,6,6,6,6,6,6,7,7,7,7,7,7],
    'quantity': [7,7,7,7,7,7,6,6,6,6,6,6,5,5,5,5,5,5,4,4,4,4,4,4],
    'estimated_price': [5,5,5,5,5,5,6,6,6,6,6,6,7,7,7,7,7,7,8,8,8,8,8,8],
    'estimated_quantity': [6,6,6,6,6,6,5,5,5,5,5,5,4,4,4,4,4,4,3,3,3,3,3,3],
    'key_product': [0.9,0.8,0.7,0.8,0.9,0.8,0.7,0.8,0.9,0.8,0.7,0.8,0.9,0.8,0.7,0.8,0.9,0.8,0.7,0.8,0.9,0.8,0.7,0.8],
    'initial_guess': [10,10,10,10,10,10,20,20,20,20,20,20,30,30,30,30,30,30,40,40,40,40,40,40]
}).lazy()
hourly_data = hourly_data.with_columns(pl.datetime_range(datetime(2024,1,1), datetime(2024,1,1,23), '1h').alias('hour'))
hourly_data = hourly_data.with_columns(pl.col('hour').min().over('period').alias('period_start'))

# 预计算每个周期的聚合参数
period_params = hourly_data.group_by('period_start', maintain_order=True).agg(
    pl.col('price').median().alias('price_med'),
    pl.col('quantity').median().alias('qty_med'),
    pl.col('estimated_quantity').sum().alias('est_qty_sum'),
    pl.col('estimated_quantity').mul(pl.col('estimated_price')).sum().alias('est_qty_price_sum'),
    pl.col('key_product').sum().alias('key_sum'),
    pl.col('key_product').mul(pl.col('estimated_price')).sum().alias('key_price_sum'),
    pl.col('initial_guess').median().alias('init_guess')
).collect()

步骤2:重构目标函数为纯数值运算

直接使用预计算的参数,无需再操作Polars数据:

def fast_target(x, params):
    price_med, qty_med, est_qty_sum, est_qty_price_sum, key_sum, key_price_sum = params
    numerator = abs((price_med * qty_med) - (est_qty_price_sum + x * est_qty_sum))
    denominator = abs(key_price_sum + x * key_sum)
    return abs(numerator / denominator - 1)

步骤3:批量执行优化

遍历每个周期的参数,调用scipy优化:

results = []
for row in period_params.rows(named=True):
    params = (row['price_med'], row['qty_med'], row['est_qty_sum'], row['est_qty_price_sum'], row['key_sum'], row['key_price_sum'])
    res = scipy.optimize.minimize(fast_target, row['init_guess'], args=(params,))
    results.append({'period_start': row['period_start'], 'x_values': res.x[0]})

# 转换为Polars DataFrame
results_df = pl.DataFrame(results).sort('period_start')
print(results_df)

优化效果说明

  • 仅做一次Polars聚合计算,后续全为纯数值运算,避免了重复查询的开销,速度提升显著。
  • 目标函数逻辑与原代码完全一致,输出结果与期望完全匹配。

内容的提问来源于stack exchange,提问作者thanksgivingmeal

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 22:47:32