如何更快找到多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
相关产品推荐
相关产品推荐

