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

如何在neat-python的eval_genome函数中获取世代编号?

问题

我正在使用neat-python库进行神经网络相关实验,当前的eval_genome函数实现如下:

def eval_genome(genome, config):
    pheno = neat.nn.FeedForwardNetwork.create(genome, config)
    data = open("RANDOM FILE", "R") # 随机生成的文件
    return pheno.activate(data)[0]

def train():
    config = neat.Config(neat.DefaultGenome, neat.DefaultReproduction, neat.DefaultSpeciesSet, neat.DefaultStagnation, os.path.join(os.path.dirname(__file__), "neat_config"))
    pop = neat.Population(config)

    pop.add_reporter(neat.StdOutReporter(True))
    pop.add_reporter(neat.Checkpointer(1, 2 ** 64, "checkpoints/checkpoint-"))

    pe = neat.ParallelEvaluator(multiprocessing.cpu_count(), eval_genome)
    winner = pop.run(pe.evaluate, 300)

直接在eval_genome内生成随机数据集会导致同一世代的神经网络无法使用统一基准数据集,部分网络可能因数据集简单获得不公平优势。我希望能在eval_genome函数中获取世代编号,实现如下逻辑:

generations_datasets = {}

def eval_genome(genome, config, generation_number):

  if not (generation_number in generations_datasets):
    generations_datasets[generation_number] = # 生成随机数据集

  pheno = neat.nn.FeedForwardNetwork.create(genome, config)
  data = generations_datasets[generation_number]
  return pheno.activate(data)[0]

请问是否有可行的实现方法?

解决方案

方法1:自定义Reporter跟踪世代编号

neat-python的Population.run()会在每世代启动时触发Reporter的start_generation方法,我们可以通过自定义Reporter记录当前世代号,再结合全局变量让eval_genome获取该编号,提前为每个世代生成统一数据集。

示例代码:

import neat
import os
import multiprocessing

current_generation = 0
generations_datasets = {}

class GenerationTracker(neat.reporting.BaseReporter):
    def start_generation(self, gen):
        global current_generation
        current_generation = gen
        # 提前生成当前世代的数据集
        if gen not in generations_datasets:
            generations_datasets[gen] = generate_random_dataset()

def generate_random_dataset():
    # 替换为你的数据集生成逻辑,返回可直接使用的数据而非文件对象
    import random
    return [random.random() for _ in range(10)]

def eval_genome(genome, config):
    global current_generation
    data = generations_datasets[current_generation]
    pheno = neat.nn.FeedForwardNetwork.create(genome, config)
    return pheno.activate(data)[0]

def train():
    config_path = os.path.join(os.path.dirname(__file__), "neat_config")
    config = neat.Config(neat.DefaultGenome, neat.DefaultReproduction,
                         neat.DefaultSpeciesSet, neat.DefaultStagnation, config_path)
    pop = neat.Population(config)

    # 添加自定义世代跟踪器
    pop.add_reporter(GenerationTracker())
    pop.add_reporter(neat.StdOutReporter(True))
    pop.add_reporter(neat.Checkpointer(1, 2**64, "checkpoints/checkpoint-"))

    pe = neat.ParallelEvaluator(multiprocessing.cpu_count(), eval_genome)
    winner = pop.run(pe.evaluate, 300)

注意:多进程场景下全局变量无法在子进程共享,此时需将数据集序列化(比如存为临时文件、用multiprocessing.Manager管理字典),确保子进程能读取到对应世代的数据集。单进程评估可直接使用上述代码。

方法2:手动控制世代循环+闭包传递编号

放弃Population.run()的自动循环,手动控制每世代流程,通过闭包将世代编号传入eval_genome,确保同一世代的所有个体使用相同数据集。

示例代码:

import neat
import os
import multiprocessing

generations_datasets = {}

def create_eval_func(generation_num):
    def eval_genome(genome, config):
        if generation_num not in generations_datasets:
            generations_datasets[generation_num] = generate_random_dataset()
        data = generations_datasets[generation_num]
        pheno = neat.nn.FeedForwardNetwork.create(genome, config)
        return pheno.activate(data)[0]
    return eval_genome

def generate_random_dataset():
    import random
    return [random.random() for _ in range(10)]

def train():
    config_path = os.path.join(os.path.dirname(__file__), "neat_config")
    config = neat.Config(neat.DefaultGenome, neat.DefaultReproduction,
                         neat.DefaultSpeciesSet, neat.DefaultStagnation, config_path)
    pop = neat.Population(config)

    pop.add_reporter(neat.StdOutReporter(True))
    pop.add_reporter(neat.Checkpointer(1, 2**64, "checkpoints/checkpoint-"))

    # 手动循环300世代
    for gen in range(300):
        # 创建绑定当前世代编号的评估函数
        eval_func = create_eval_func(gen)
        pe = neat.ParallelEvaluator(multiprocessing.cpu_count(), eval_func)
        # 评估所有个体
        fitness_map = pe.evaluate(pop.population.items(), config)
        # 更新个体适应度
        for genome_id, genome in pop.population.items():
            genome.fitness = fitness_map[genome_id]
        # 推进种群进化
        pop.reproduce()
        pop.species.speciate(config, pop.population, pop.generation)
        pop.generation += 1
        # 触发Reporter的世代结束事件
        for reporter in pop.reporters.reporters:
            if hasattr(reporter, 'end_generation'):
                reporter.end_generation(config, pop.population, pop.species)
    
    # 获取最优个体
    winner = max(pop.population.values(), key=lambda g: g.fitness)

这种方式完全掌控世代流程,数据集生成逻辑清晰,多进程场景下只需确保数据集可被子进程访问即可。

方法3:用functools.partial绑定世代参数

结合自定义Reporter跟踪世代,用functools.partial给eval_genome绑定世代编号参数,既保留Population.run()的便捷性,又能传递额外参数。

示例代码:

import neat
import os
import multiprocessing
from functools import partial

generations_datasets = {}
current_eval_func = None

class GenerationReporter(neat.reporting.BaseReporter):
    def start_generation(self, gen):
        global current_eval_func
        # 绑定世代编号到eval_genome
        current_eval_func = partial(eval_genome, generation_number=gen)
        # 预生成当前世代数据集
        if gen not in generations_datasets:
            generations_datasets[gen] = generate_random_dataset()

def eval_genome(genome, config, generation_number):
    data = generations_datasets[generation_number]
    pheno = neat.nn.FeedForwardNetwork.create(genome, config)
    return pheno.activate(data)[0]

def generate_random_dataset():
    import random
    return [random.random() for _ in range(10)]

def custom_evaluate(population, config):
    global current_eval_func
    pe = neat.ParallelEvaluator(multiprocessing.cpu_count(), current_eval_func)
    return pe.evaluate(population, config)

def train():
    config_path = os.path.join(os.path.dirname(__file__), "neat_config")
    config = neat.Config(neat.DefaultGenome, neat.DefaultReproduction,
                         neat.DefaultSpeciesSet, neat.DefaultStagnation, config_path)
    pop = neat.Population(config)

    pop.add_reporter(GenerationReporter())
    pop.add_reporter(neat.StdOutReporter(True))
    pop.add_reporter(neat.Checkpointer(1, 2**64, "checkpoints/checkpoint-"))

    # 使用自定义评估函数启动训练
    winner = pop.run(custom_evaluate, 300)

该方法通过partial动态绑定世代参数,无需手动循环,同时保证同一世代的所有个体使用相同数据集。


内容的提问来源于stack exchange,提问作者Ank i zle

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 19:00:23