如何用Polars无循环模拟多节点服务队列?
问题
我正在模拟节点数可配置的先进先出(FIFO)服务队列,需要创建自定义函数或Polars表达式(支持Python/Rust插件)来计算每个对象的服务开始与结束时间。目前通过遍历DataFrame行实现,想寻求更高效、符合Polars惯用写法的方案。
输入DataFrame示例:
import polars as pl df = pl.DataFrame({ 'id': [1,2,3,4], 'servicing_time_requirement': [30, 5, 30, 5], 'arrival_time': [0, 15, 16, 17], })
期望输出示例
- 节点数=5(无等待,所有任务立即开始):
| id | servicing_time_requirement | arrival_time | service_start_time | service_end_time |
|---|---|---|---|---|
| 1 | 30 | 0 | 0 | 30 |
| 2 | 5 | 15 | 15 | 20 |
| 3 | 30 | 16 | 16 | 46 |
| 4 | 5 | 17 | 17 | 22 |
- 节点数=2(需等待空闲节点):
| id | servicing_time_requirement | arrival_time | service_start_time | service_end_time |
|---|---|---|---|---|
| 1 | 30 | 0 | 0 | 30 |
| 2 | 5 | 15 | 15 | 20 |
| 3 | 30 | 16 | 20 | 50 |
| 4 | 5 | 17 | 30 | 35 |
当前逐行遍历的Python实现(注:原代码存在列名错误,已修正):
nodes = 2 arrival_times = df.get_column("arrival_time") servicing_end_times = pl.Series([None] * len(df)) servicing_time_requirements = df.get_column("servicing_time_requirement") for i in range(servicing_end_times.len()): if servicing_end_times[i] is not None: continue next_done = servicing_end_times.filter( servicing_end_times.is_not_null() ).rank(method='ordinal', descending=True).eq(nodes) if next_done.len() == 0: next_done = arrival_times[i] else: next_done = max(servicing_end_times.filter(next_done)[0], arrival_times[i]) servicing_end_times[i] = next_done + servicing_time_requirements[i]
解决方案
方法1:Polars表达式结合堆(Python实现)
利用Polars的map_batches方法,结合Python的最小堆结构实现高效节点调度,避免逐行遍历的低效问题。核心逻辑是维护一个记录节点空闲时间的堆,每个任务分配给最早空闲的节点。
import polars as pl import heapq def calculate_service_times(servicing_times, arrival_times, nodes): heap = [] start_times = [] end_times = [] # 按到达时间遍历任务(需确保DataFrame已排序) for st, at in zip(servicing_times, arrival_times): if len(heap) < nodes: # 有空闲节点,直接开始服务 start = at else: # 取出最早空闲的节点时间 earliest_end = heapq.heappop(heap) start = max(at, earliest_end) end = start + st start_times.append(start) end_times.append(end) # 将当前任务的结束时间放回堆 heapq.heappush(heap, end) return start_times, end_times # 先按到达时间排序(FIFO队列必须保证顺序) df_sorted = df.sort("arrival_time") nodes = 2 result_df = df_sorted.with_columns( pl.struct(["servicing_time_requirement", "arrival_time"]) .map_batches( lambda batch: pl.DataFrame( calculate_service_times( batch["servicing_time_requirement"].to_list(), batch["arrival_time"].to_list(), nodes ), columns=["service_start_time", "service_end_time"] ) ).alias("service_times") ).unnest("service_times") print(result_df)
方法2:Rust插件(超大数据量场景)
如果处理超大规模数据集,用Rust编写Polars自定义表达式插件能获得更高性能。核心逻辑与Python堆实现一致,但Rust底层操作更高效。
以下是Rust插件的核心代码:
use polars::prelude::*; use std::collections::BinaryHeap; use std::cmp::Reverse; #[polars_expr(output_type=Struct)] fn calculate_service_times(inputs: &[Series]) -> PolarsResult<Series> { // 解析输入序列:服务时长、到达时间、节点数 let servicing_times = inputs[0].i64()?; let arrival_times = inputs[1].i64()?; let nodes = inputs[2].i64()?[0] as usize; let mut heap = BinaryHeap::new(); let mut start_times = Vec::with_capacity(servicing_times.len()); let mut end_times = Vec::with_capacity(servicing_times.len()); for (st, at) in servicing_times.iter().zip(arrival_times.iter()) { let st = st.unwrap(); let at = at.unwrap(); if heap.len() < nodes { let start = at; let end = start + st; start_times.push(start); end_times.push(end); heap.push(Reverse(end)); } else { let earliest_end = heap.pop().unwrap().0; let start = std::cmp::max(at, earliest_end); let end = start + st; start_times.push(start); end_times.push(end); heap.push(Reverse(end)); } } // 构造输出的Struct序列 let start_series = Series::new("service_start_time", start_times); let end_series = Series::new("service_end_time", end_times); let struct_series = StructSeries::new(vec![start_series, end_series])?; Ok(struct_series.into_series()) }
编译插件后,即可在Python中通过Polars调用该自定义表达式,处理百万级以上数据时性能远超Python实现。
内容的提问来源于stack exchange,提问作者mb7155
相关产品推荐
相关产品推荐

