diff --git a/pygad/utils/mutation.py b/pygad/utils/mutation.py index 66d72e29..a7d10258 100644 --- a/pygad/utils/mutation.py +++ b/pygad/utils/mutation.py @@ -362,9 +362,8 @@ def polynomial_mutation(self, offspring): def swap_mutation(self, offspring): """ - Swap the values of two genes inside each offspring. One gene is - picked at random from the first half of the chromosome; the - other is its mirror in the second half. + Swap the values of two genes inside each offspring. The two + genes are 2 different genes picked at random. Parameters ---------- @@ -378,8 +377,7 @@ def swap_mutation(self, offspring): """ for idx in range(offspring.shape[0]): - mutation_gene1 = numpy.random.randint(low=0, high=offspring.shape[1]/2, size=1)[0] - mutation_gene2 = mutation_gene1 + int(offspring.shape[1]/2) + mutation_gene1, mutation_gene2 = numpy.random.choice(offspring.shape[1], size=2, replace=False) temp = offspring[idx, mutation_gene1] offspring[idx, mutation_gene1] = offspring[idx, mutation_gene2] diff --git a/tests/test_crossover_mutation.py b/tests/test_crossover_mutation.py index acc38942..7f51fdc0 100644 --- a/tests/test_crossover_mutation.py +++ b/tests/test_crossover_mutation.py @@ -241,6 +241,26 @@ def test_random_mutation_manual_call4(): for value in comp_sorted: assert value in value_space +def test_swap_mutation_manual_call(): + # Any 2 different genes can be swapped. + num_genes = 6 + result, ga_instance = output_crossover_mutation(gene_type=int, + num_genes=num_genes, + mutation_type="swap") + + temp_offspring = numpy.array([list(range(num_genes))] * 1000) + offspring = ga_instance.swap_mutation(offspring=temp_offspring.copy()) + + swapped_pairs = set() + for solution in offspring: + changed = numpy.flatnonzero(solution != numpy.arange(num_genes)) + # Exactly 2 genes exchange their values. + assert len(changed) == 2 + assert solution[changed[0]] == changed[1] and solution[changed[1]] == changed[0] + swapped_pairs.add(tuple(changed)) + + assert len(swapped_pairs) == num_genes * (num_genes - 1) // 2 + if __name__ == "__main__": #### Single-objective print() @@ -285,3 +305,6 @@ def test_random_mutation_manual_call4(): test_random_mutation_manual_call4() print() + + test_swap_mutation_manual_call() + print()