如何用Pulp将TSP扩展为MTSP?现有代码路径求解错误求助
多旅行商问题(MTSP)Pulp代码修复求助
我们已研究过旅行商问题(TSP),现需扩展为多旅行商问题(MTSP)。使用Pulp编写代码并添加自定义逻辑后,无法正确生成最优旅行路径。
示例数据与预期结果
成本矩阵
cost_matrix = [[ 0, 1, 3, 4], [ 1, 0, 2, 3 ], [ 3, 2, 0, 4 ], [ 4, 3, 4, 0]]
设n = len(cost_matrix),当旅行商数量k=3时,预期结果为:
- SP_1的路径:0 => 1 => 0
- SP_2的路径:0 => 2 => 3 => 0
现有代码
# create encoding variables bin_vars = [ # add a binary variable x_{ij} if i not = j else simply add None [ LpVariable(f'x_{i}_{j}', cat='Binary') if i != j else None for j in range(n)] for i in range(n) ] time_stamps = [LpVariable(f't_{j}', lowBound=0, upBound=n, cat='Continuous') for j in range(1, n)] # create add the objective function objective_function = lpSum( [ lpSum([xij*cj if xij != None else 0 for (xij, cj) in zip(brow, crow) ]) for (brow, crow) in zip(bin_vars, cost_matrix)] ) prob += objective_function # add constraints for i in range(n): # Exactly one leaving variable prob += lpSum([xj for xj in bin_vars[i] if xj != None]) == 1 # Exactly one entering prob += lpSum([bin_vars[j][i] for j in range(n) if j != i]) == 1 # add timestamp constraints for i in range(1,n): for j in range(1, n): if i == j: continue xij = bin_vars[i][j] ti = time_stamps[i-1] tj = time_stamps[j -1] prob += tj >= ti + xij - (1-xij)*(n+1) # Binary variables to ensure each node is visited by a salesperson visit_vars = [LpVariable(f'u_{i}', cat='Binary') for i in range(1, n)] # Salespersons constraints prob += lpSum([bin_vars[0][j] for j in range(1, n)]) == k prob += lpSum([bin_vars[i][0] for i in range(1, n)]) == k for i in range(1, n): prob += lpSum([bin_vars[i][j] for j in range(n) if j != i]) == visit_vars[i - 1] prob += lpSum([bin_vars[j][i] for j in range(n) if j != i]) == visit_vars[i - 1] # Done: solve the problem status = prob.solve(PULP_CBC_CMD(msg=False))
问题分析与修复方案
1. 节点出入度约束冲突
原代码对所有节点(包括起点0)设置了「恰好1条出边和1条入边」的约束,但MTSP中起点0需要有k条出边和k条入边(对应k个旅行商),其他节点保持1出1入。需修改约束:
# 非起点节点出入度为1 for i in range(1, n): prob += lpSum([xj for xj in bin_vars[i] if xj != None]) == 1 prob += lpSum([bin_vars[j][i] for j in range(n) if j != i]) == 1 # 起点0的出入度设为k prob += lpSum([xj for xj in bin_vars[0] if xj != None]) == k prob += lpSum([bin_vars[j][0] for j in range(n) if j != 0]) == k
2. 子回路消除逻辑错误
原时间戳约束仅覆盖非起点节点,且逻辑不符合MTSP子回路消除要求。替换为标准约束:
# 重新定义时间戳变量(包含起点0) time_stamps = [LpVariable(f't_{j}', lowBound=0, upBound=n, cat='Continuous') for j in range(n)] # 子回路消除约束 for i in range(n): for j in range(1, n): if i == j: continue xij = bin_vars[i][j] prob += time_stamps[j] >= time_stamps[i] + xij - (n - 1) * (1 - xij)
3. 冗余变量移除
visit_vars变量完全多余,非起点节点的出入度约束已保证每个节点被恰好访问一次,直接删除相关代码块。
修复后的完整代码
import pulp cost_matrix = [[0, 1, 3, 4], [1, 0, 2, 3], [3, 2, 0, 4], [4, 3, 4, 0]] n = len(cost_matrix) k = 3 # 旅行商数量 # 初始化问题 prob = pulp.LpProblem("MTSP", pulp.LpMinimize) # 创建路径变量x_ij:i到j的路径是否存在 bin_vars = [ [pulp.LpVariable(f'x_{i}_{j}', cat='Binary') if i != j else None for j in range(n)] for i in range(n) ] # 创建时间戳变量用于子回路消除 time_stamps = [pulp.LpVariable(f't_{j}', lowBound=0, upBound=n, cat='Continuous') for j in range(n)] # 目标函数:总路径成本最小 objective_function = pulp.lpSum( [pulp.lpSum([xij * cj for xij, cj in zip(brow, crow) if xij is not None]) for brow, crow in zip(bin_vars, cost_matrix)] ) prob += objective_function # 约束:非起点节点出入度为1 for i in range(1, n): prob += pulp.lpSum([xj for xj in bin_vars[i] if xj is not None]) == 1 prob += pulp.lpSum([bin_vars[j][i] for j in range(n) if j != i]) == 1 # 约束:起点0的出入度为k prob += pulp.lpSum([xj for xj in bin_vars[0] if xj is not None]) == k prob += pulp.lpSum([bin_vars[j][0] for j in range(n) if j != 0]) == k # 约束:子回路消除 for i in range(n): for j in range(1, n): if i == j: continue xij = bin_vars[i][j] prob += time_stamps[j] >= time_stamps[i] + xij - (n - 1) * (1 - xij) # 求解问题 status = prob.solve(pulp.PULP_CBC_CMD(msg=False)) # 提取并输出路径 paths = [] visited = [False] * n start_node = 0 for _ in range(k): # 找到当前旅行商的起始分支 current = start_node if visited[current]: for j in range(n): if not visited[j] and pulp.value(bin_vars[start_node][j]) == 1: current = j break path = [current] visited[current] = True # 遍历路径直到返回起点 while True: next_node = None for j in range(n): if current != j and pulp.value(bin_vars[current][j]) == 1: next_node = j break if next_node == start_node: path.append(next_node) break path.append(next_node) visited[next_node] = True current = next_node paths.append(path) # 打印结果 for idx, path in enumerate(paths, 1): print(f"SP_{idx}的路径:{' => '.join(map(str, path))}")
内容的提问来源于stack exchange,提问作者Sys
相关产品推荐
相关产品推荐

