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

PyGAD设置allow_duplicate_genes=False求解TSP仍出现重复基因问题

PyGAD求解TSP时allow_duplicate_genes参数失效排查方案

问题描述

  • 初始化PyGAD遗传算法实例时已显式设置allow_duplicate_genes=False,传入的初始种群initial_population为numpy数组格式,仍出现迭代后最优解存在大量重复基因的问题,无法满足TSP路径节点不重复的约束。
  • 所用初始种群样例:
[[ 1 12 26 19 22 20  6 15 17 23 21  7 28  5 13 14 16  2 24  4  3 10  9  8  18 25 27 11]
 [ 2 17 23 27 22 12 20 21 24 25 13  5 10  4  9 26  7  1 11  3 15 18 16 14 8 19 28  6]
 [ 3 23 12  7  2 11 15 13 19 26 21 14  9  5 24 20 25  1  8 16 22 28 27 10 4  6 18 17]
 [ 4 19  2 25 21 13  98 18 28  7 27 20 11 23 22 14  1 10 16 12  5 26 24 17  3 15  6]
 [ 5  9 19  7 22 10 11 13  1 25  6 17  8 12  2 24 28 20 26  4 15 14 18 23 21 27  3 16]]
  • 输出的含重复值的无效结果样例:
    [ 6 8 13 1 19 10 6 23 18 22 5 3 21 11 6 16 28 1 4 10 6 25 7 22 5 3 21 11]
  • 核心实现代码:
import pygad
import numpy as np
import copy


def fitness_func(solution, solution_idx):
        distance_treshold=np.load('distance.npy')
        function_inputs=distance_simple(distance_treshold)
        a1=treshold(function_inputs)
        f=0
        for i in range(len(solution)):
         if i == 0:
            f+= distance_treshold[solution[0]][solution[i+1]]   
         else:
          try:
            f+= distance_treshold[solution[i]][solution[i+1]]
   
          except:

             f+=distance_treshold[solution[i]][solution[0]]
                    
               
        fitness_score=pow((a1)/f,2)#fitness
        return fitness_score

def treshold(solution):
    distance_treshold=np.load('distance.npy')
    f=0
    for i in range(len(solution)):
             if i == 0:
                f+= distance_treshold[solution[0]][solution[i+1]]    
             else:
                try:
                    f+= distance_treshold[solution[i]][solution[i+1]]
   
                except:

                    f+=distance_treshold[solution[i]][solution[0]] 
    return f


function_inputs=distance_simple(distance_treshold)
a1=treshold(function_inputs)
print(a1)
np.load('distance.npy')

#print(initial_pop)
initial_population=np.load('inital_generation.npy')
print(initial_population)
num_parents_mating= 2
num_generations= 30
parent_selection_type='sus'
mutation_type="swap"
keep_parents=0
mutation_num_genes=1
mutation_percent_genes=3
crossover_type="single_point"
allow_duplicate_genes=False
gene_type=int
mutation_probability=0.03

print("GA start")
ga_instance = pygad.GA(num_generations=num_generations,mutation_probability=mutation_probability,
parent_selection_type=parent_selection_type,initial_population=initial_population,
num_parents_mating=num_parents_mating, fitness_func=fitness_func,gene_type=gene_type, 
mutation_percent_genes=mutation_num_genes,mutation_num_genes=mutation_percent_genes,
mutation_type=mutation_type,allow_duplicate_genes=False)


ga_instance.run()
ga_instance.plot_fitness()
best_solution,best_solution_fitness,best_match_idx=ga_instance.best_solution()
print(best_solution)
fitness_func(best_solution,0)
print(best_solution_fitness)

根因定位

按影响优先级从高到低排列:

  • 变异参数传反:初始化GA实例时,mutation_percent_genes和mutation_num_genes两个参数的传入值完全颠倒。原本定义mutation_num_genes=1(每次变异只改1个基因)、mutation_percent_genes=3(变异比例3%),传参时互换了值,导致变异逻辑完全失控,大量产生重复基因。
  • 参数漏传:代码中提前定义了crossover_type、keep_parents参数,但构造GA实例时根本没有传入这两个参数,实际运行用的是PyGAD默认值,和预期配置完全不符。
  • 交叉算子选型错误:原本计划用的single_point(单点交叉)是为二进制/实数编码设计的,天生不适合TSP这类排列编码问题——单点交叉通过切割两个父代片段直接拼接后代,90%以上概率会产生重复基因,即使开启allow_duplicate_genes=False,常规交叉算子的去重校验也无法完全修复排列破坏问题。
  • 初始种群存在无效值:初始种群第四行个体出现了超出节点范围的98(总节点仅28个,编号范围1~28),属于非法基因。
  • 适应度函数逻辑缺陷:用try-except直接吞掉所有索引越界错误,越界的非法个体也能计算出适应度,获得进入下一代的资格;同时每次计算适应度都重复从磁盘加载distance.npy,性能极差。
  • 缺失基因空间约束:没有传入gene_space参数明确基因取值范围,变异、交叉环节生成新基因时无边界约束,容易产生超出节点范围的非法值。

修复方案

按以下步骤逐一修改即可解决重复基因问题:

  • 第一步:修正参数传递错误,把颠倒的变异参数改回正确值,同时把漏传的crossover_type、keep_parents参数补全,替换交叉算子为排列编码适配的类型。TSP场景下建议交叉算子选用"two_points"搭配swap变异,同时明确传入gene_space为节点编号范围。
  • 第二步:清理初始种群,删除所有超出节点编号范围的非法值(比如样例中的98),保证初始种群所有个体都是1~N的不重复整数排列(N为节点总数)。
  • 第三步:重写适应度函数,把距离矩阵加载逻辑移到函数外部,删除吞错误的try-except,对越界、重复的非法个体直接返回极小的适应度值,避免非法个体存活。
  • 第四步:调整keep_parents参数为≥1,保证每代最优的父代个体可以直接保留,避免交叉变异完全破坏优质解。

修复后的核心代码参考:

import pygad
import numpy as np

# 提前加载距离矩阵,避免重复读磁盘
distance_treshold = np.load('distance.npy')
node_count = distance_treshold.shape[0]
# 定义基因空间:节点编号从0开始用range(node_count),从1开始用range(1, node_count+1)
gene_space = list(range(node_count))

def fitness_func(solution, solution_idx):
    # 提前校验个体合法性,有重复直接给极低分
    if len(np.unique(solution)) != len(solution):
        return 1e-8
    f = 0
    for i in range(len(solution)):
        if i == len(solution)-1:
            # 最后一个节点连回起点
            f += distance_treshold[solution[i]][solution[0]]
        else:
            f += distance_treshold[solution[i]][solution[i+1]]
    # 距离越短适应度越高
    fitness_score = 1/(f**2)
    return fitness_score

# 加载并清洗初始种群
initial_population = np.load('inital_generation.npy')
valid_initial_pop = []
for ind in initial_population:
    if len(np.unique(ind)) == len(ind) and np.max(ind) < node_count and np.min(ind) >=0:
        valid_initial_pop.append(ind)
initial_population = np.array(valid_initial_pop)

# GA参数配置
num_parents_mating= 2
num_generations= 100
parent_selection_type='sus'
mutation_type="swap"
keep_parents=2
mutation_num_genes=1
mutation_percent_genes=3
crossover_type="two_points"
allow_duplicate_genes=False
gene_type=int
mutation_probability=0.05

ga_instance = pygad.GA(
    num_generations=num_generations,
    mutation_probability=mutation_probability,
    parent_selection_type=parent_selection_type,
    initial_population=initial_population,
    num_parents_mating=num_parents_mating,
    fitness_func=fitness_func,
    gene_type=gene_type,
    gene_space=gene_space,
    mutation_percent_genes=mutation_percent_genes,
    mutation_num_genes=mutation_num_genes,
    mutation_type=mutation_type,
    allow_duplicate_genes=allow_duplicate_genes,
    crossover_type=crossover_type,
    keep_parents=keep_parents
)

ga_instance.run()
ga_instance.plot_fitness()
best_solution,best_solution_fitness,best_match_idx=ga_instance.best_solution()
print("最优路径:", best_solution)
print("最优路径适应度:", best_solution_fitness)
# 校验最优解合法性
print("最优解是否无重复节点:", len(np.unique(best_solution)) == len(best_solution))

验证要点

运行修复后的代码时,最后打印的最优解是否无重复节点会返回True,不会再出现同一路径重复访问同一节点的问题。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 06:33:21