如何在PySpark中求解含多任务边界的指派问题(匈牙利算法/linear_sum_assignment)
Spark 3.1+ 带特殊约束指派问题最优实现方案
你的场景N最大值仅为10,单分组数据量极小,基于applyInPandas改造是成本最低、性能足够的方案,只需要在原有逻辑基础上增加多任务约束前置处理即可,完整实现如下:
核心处理逻辑
每个分组内按固定顺序处理约束:
- 先判断当前分组是否存在
isMultiTask=True的任务,不存在则直接走普通线性分配逻辑 - 若存在多任务,先筛选所有执行该任务成本低于阈值x的工作人员,这部分人直接分配多任务,从待分配人员列表中移除
- 若有至少1人自动分配了多任务,该多任务直接标记为已占用,从待分配任务列表移除;如果没人满足成本阈值,多任务保留和普通任务一起参与分配
- 剩余工作人员和待分配任务调用
linear_sum_assignment求解最优分配 - 合并自动分配结果和优化分配结果返回
完整实现代码
依赖导入与基础索引生成
import pandas as pd import numpy as np from scipy.optimize import linear_sum_assignment from pyspark.sql import SparkSession from pyspark.sql.window import Window from pyspark.sql.functions import dense_rank from pyspark.sql.types import StructType, StringType, IntegerType, DoubleType, BooleanType # SparkSession初始化按需调整 spark = SparkSession.builder.appName("AssignmentProblem").getOrCreate() # 你的样例数据生成逻辑此处省略,和你提供的代码一致 # 基础索引生成逻辑和原有方案一致 worker_order_window = Window.partitionBy("date", "locationId").orderBy("workerId") task_order_window = Window.partitionBy("date", "locationId").orderBy("taskId") df = df.withColumn("worker_idx", dense_rank().over(worker_order_window) - 1) df = df.withColumn("task_idx", dense_rank().over(task_order_window) - 1)
改造后的Pandas UDF
# 多任务成本阈值可根据业务调整,也可作为参数动态传入UDF MULTI_TASK_THRESHOLD = 1 def assignment_with_multi_task(pandas_df: pd.DataFrame) -> pd.DataFrame: result = [] # 提取当前分组多任务信息 multi_task_df = pandas_df[pandas_df["isMultiTask"] == True] # 场景1:无多任务,直接走普通线性分配 if len(multi_task_df) == 0: worker_list = pandas_df["workerId"].unique() task_list = pandas_df["taskId"].unique() n_worker, n_task = len(worker_list), len(task_list) N = max(n_worker, n_task) # 用远大于最大成本10的值填充空位置,不影响最优解 cost_mat = np.full((N, N), 100.0) worker_map = {w:i for i,w in enumerate(worker_list)} task_map = {t:i for i,t in enumerate(task_list)} # 填充成本矩阵 for _, row in pandas_df.iterrows(): w_idx = worker_map[row["workerId"]] t_idx = task_map[row["taskId"]] cost_mat[w_idx][t_idx] = row["cost"] # 求解并组装结果 rids, cids = linear_sum_assignment(cost_mat) assign_map = {} for r, c in zip(rids, cids): if r < n_worker and c < n_task and cost_mat[r][c] < 100: assign_map[worker_list[r]] = task_list[c] for _, row in pandas_df.iterrows(): assign_task = assign_map.get(row["workerId"], -1) row["task_assignment"] = assign_task row["isAssigned"] = assign_task == row["taskId"] result.append(row) return pd.DataFrame(result) # 场景2:存在多任务,先处理自动分配逻辑 multi_task_id = multi_task_df["taskId"].iloc[0] # 筛选满足阈值的自动分配人员 auto_assign_workers = multi_task_df[multi_task_df["cost"] < MULTI_TASK_THRESHOLD]["workerId"].unique() # 先写入自动分配结果 for worker_id in auto_assign_workers: worker_rows = pandas_df[pandas_df["workerId"] == worker_id] for _, row in worker_rows.iterrows(): row["task_assignment"] = multi_task_id row["isAssigned"] = row["taskId"] == multi_task_id result.append(row) # 生成剩余待分配的工人和任务列表 remain_workers = pandas_df[~pandas_df["workerId"].isin(auto_assign_workers)]["workerId"].unique() if len(auto_assign_workers) > 0: # 有自动分配的人,多任务不再参与后续分配 remain_tasks = pandas_df[pandas_df["isMultiTask"] == False]["taskId"].unique() else: # 没人自动分配,多任务和普通任务一起参与分配 remain_tasks = pandas_df["taskId"].unique() # 无剩余待分配资源,直接补全未分配人员结果返回 if len(remain_workers) == 0 or len(remain_tasks) == 0: for worker_id in remain_workers: worker_rows = pandas_df[pandas_df["workerId"] == worker_id] for _, row in worker_rows.iterrows(): row["task_assignment"] = -1 row["isAssigned"] = False result.append(row) return pd.DataFrame(result) # 剩余资源走线性分配 n_worker, n_task = len(remain_workers), len(remain_tasks) N = max(n_worker, n_task) cost_mat = np.full((N, N), 100.0) worker_map = {w:i for i,w in enumerate(remain_workers)} task_map = {t:i for i,t in enumerate(remain_tasks)} for _, row in pandas_df.iterrows(): if row["workerId"] in worker_map and row["taskId"] in task_map: w_idx = worker_map[row["workerId"]] t_idx = task_map[row["taskId"]] cost_mat[w_idx][t_idx] = row["cost"] # 求解并组装结果 rids, cids = linear_sum_assignment(cost_mat) assign_map = {} for r, c in zip(rids, cids): if r < n_worker and c < n_task and cost_mat[r][c] < 100: assign_map[remain_workers[r]] = remain_tasks[c] for worker_id in remain_workers: assign_task = assign_map.get(worker_id, -1) worker_rows = pandas_df[pandas_df["workerId"] == worker_id] for _, row in worker_rows.iterrows(): row["task_assignment"] = assign_task row["isAssigned"] = assign_task == row["taskId"] result.append(row) return pd.DataFrame(result)
调用执行
# 定义输出Schema output_schema = StructType() \ .add("date", StringType()) \ .add("locationId", StringType()) \ .add("workerId", IntegerType()) \ .add("taskId", IntegerType()) \ .add("cost", DoubleType()) \ .add("isMultiTask", BooleanType()) \ .add("worker_idx", IntegerType()) \ .add("task_idx", IntegerType()) \ .add("task_assignment", IntegerType()) \ .add("isAssigned", BooleanType()) # 分组求解 result_df = df.groupBy("date", "locationId").applyInPandas(assignment_with_multi_task, output_schema) result_df.show()
方案说明
- 完全覆盖所有多任务约束场景,包括0人、1人、多人满足阈值的边界情况
- 单分组计算量极小,
applyInPandas性能完全满足业务需求,无需切换其他实现方案 - 未分配人员的
task_assignment标记为-1,可直接过滤使用
内容的提问来源于stack exchange,提问作者Lauren Leder
相关产品推荐
相关产品推荐

