Evolution & control · E10 · Implementation

Evolve a population of networks with one vectorized evaluator

JAX handles the population axis, Flax defines the policy and EvoJAX updates a search distribution. The interfaces—not just the optimizer—make the experiment work.

Flax LinenJAX vmapEvoJAX PGPE
Population search moves a distribution over policy weights. The landscape is illustrative, not a measured fitness surface.
Figure 1. Search a population, not a gradient. Population search moves a distribution over policy weights. The landscape is illustrative, not a measured fitness surface. Schematic parameter-space landscape. Original vector illustration.

Follow the information

From input to outcome

Each vector is decoded into the same policy architecture. The evaluator returns fitness to PGPE, which changes the search distribution; the dashed loop is an evolutionary update, not backpropagation.

Each vector is decoded into the same policy architecture. The evaluator returns fitness to PGPE, which changes the search distribution; the dashed loop is an evolutionary update, not backpropagation.
Figure 2. Information flow. Solid arrows carry observations, tensors or artifacts; other routes are explicitly labelled. Signal shapes, matrices and network icons are schematic, not measured samples or literal neuron counts. Open full-size SVG ↗ On narrow screens, scroll the diagram horizontally.

Read this alongside Figure 1: Population search moves a distribution over policy weights. The landscape is illustrative, not a measured fitness surface. The module map and layer-level figures below expand the operations in this route.

Evolve a population of networks with one vectorized evaluator: architectureParameter population: P × parameter_count → Unflatten parameters: Flax parameter trees → Policy batch: 301 → 128 → 64 → 2 → Fitness evaluation: Same observations per candidate → PGPE update: Scores → next population. A high-level module map; comparison branches and training details are explained in the article.EVOLUTION & CONTROL / E10 / MODULE MAP01 INPUTParameter populationP × parameter_count02 MODULEUnflatten parametersFlax parameter trees03 MODULEPolicy batch301 → 128 → 64 → 204 MODULEFitness evaluationSame observations per candidate05 OUTPUTPGPE updateScores → next population
Source-grounded module map. Boxes summarize operations, not individual neurons; comparison arms and training paths are detailed below. On a small screen, scroll the diagram horizontally.
Parameter population — P × parameter_count

The architecture in context

The system we are building

Neuroevolution evaluates many candidate parameter vectors instead of differentiating a reward through a simulator. This notebook explores several implementations; the selected PGPE path combines a Flax policy with a vectorized task evaluator. JAX vmap maps one individual evaluator over the population axis, keeping the meaning of an individual evaluation explicit.

Who does what in the stack

Flax Linen
Defines the MLP and its parameter tree.
JAX vmap
Adds an explicit population evaluation axis.
EvoJAX PGPE
Generates candidates and updates the search distribution.

Custom PolicyNetwork and VectorizedTask adapters translate between flat optimizer vectors, structured Flax parameters, observations and fitness scores. The notebook also contains earlier CMA-ES, MLX and PyTorch experiments. Their presence documents exploration; it is not a controlled cross-backend speed comparison.

Framework responsibility map. Each row maps a library or custom component to its job; rows are not a sequential inference graph.
Framework responsibility map. Each row maps a library or custom component to its job; rows are not a sequential inference graph. Open full-size SVG ↗

From module map to executable structure

Inside Flax policy optimized by PGPE

Selected network definition: input dimension configurable (default 301), hidden widths 128 and 64, two outputs.

Layer-level implementation. B denotes batch size; parameter and shape conventions are expanded in the table.
Layer-level implementation. B denotes batch size; parameter and shape conventions are expanded in the table. Open full-size SVG ↗
Layer / tensor / operation ledger
Layer or branchOutput shapeImplementation detail
Observation vectorB × dFeature order is part of the checkpoint contract.
Dense 128 + ReLUB × 128Dropout .1 only when the training flag requests it.
Dense 64 + ReLUB × 64Same activation; random-key management is explicit in Flax.
Dense 2 + softmaxB × 2Two action scores from a shared policy parameter tree.
Population evaluatorPopulation × fitnessFlatten parameters; vmap evaluates candidates; PGPE updates a search distribution.

Flax owns the parameter tree and forward pass. JAX vmap maps the same evaluation function over candidates; it does not make variable-length Python side effects magically vectorized. PGPE changes the distribution over weights using fitness evaluations. A gradient of the network’s logits is not required for that search.

The equation and the update

θi=μ+σ⊙ϵi,J(μ,σ)=Eϵ[F(θi)]\theta_i=\mu+\sigma\odot\epsilon_i,\qquad J(\mu,\sigma)=\mathbb E_{\epsilon}[F(\theta_i)]

PGPE uses rollout fitness, not backpropagation through a supervised loss. Population size, rollout count and generation budget determine evaluation work; the selected notebook contains multiple runners, not one uniquely identified benchmark.

Learning or solution path. A parameter-update path is different from the forward inference path; see text for target-network, frozen-feature and local-loss boundaries.
Learning or solution path. A parameter-update path is different from the forward inference path; see text for target-network, frozen-feature and local-loss boundaries. Open full-size SVG ↗

Implementation card / no invented benchmarks

Capacity, budget and execution evidence

Parameters / retained state
47,042 policy scalars when d=301; general count 128(d+1)+64×129+2×65.
Duration and hardware evidence
No single source-backed training duration applies to all notebook runners.
Source coordinates
E10 cell 9, network definition and population fitness evaluator
Current reproduction context
Current workstation, supplied by the author: Apple M4, 128 GB unified RAM, 40 GPU cores and 16 CPU cores. This is context for prospective reproduction, not attribution of every archived run. Python and framework versions are not fully locked for these historical sources; declarations, when available, are identified separately.

Counts above are calculated from the stated layer shapes unless identified as saved measurements. They exclude optimizer state and nontrainable buffers. No archived training was rerun for this revision.

What these design choices change

A small fixed MLP makes the flat search vector manageable. Doubling hidden width can substantially increase both the search dimension and population evaluation memory. Dropout noise and environment noise must be separated from perturbation noise or candidate rankings become needlessly variable.

Reproduction and measurement protocol

Round-trip the parameter pytree through flatten/unflatten and compare outputs exactly before running a population. Evaluate the same candidate and random seed twice. Time compilation separately from warm evaluation, and report candidates times rollout steps rather than calling a generation equivalent to a supervised epoch.

For a new run, save the resolved Python/framework versions, backend, dtype, seed, input shapes, batch size and exact source revision. Start with one batch and one update. Log training steps separately from epochs or environment steps. Do not equate the configured maximum with a completed budget or convergence.

Measure initialization/compilation, data preparation, warmed forward pass, training updates and evaluation separately. Synchronize accelerator work around timed regions using the chosen framework’s supported mechanism. Report peak process memory and framework allocation separately; parameter bytes exclude activations, gradients, optimizer state and input buffers. On a shared machine, begin with a single CPU worker and a small batch rather than claiming all available resources.

A closer look at the implementation

The code that carries the idea

The snippet returns one scalar fitness per parameter vector. The vectorization boundary is the candidate, not a time step or a training batch. Parameters are supplied as data to the policy rather than mutated inside a shared model object. This is the key to avoiding accidental state sharing across candidates.

Python · cell 9 · lines 268–288
    def evaluate_population(self, 
                           policy: EntryDetectionPolicy,
                           params_batch: jnp.ndarray) -> jnp.ndarray:
        """
        Evaluate population fitness in batch
        
        Args:
            policy: Policy network
            params_batch: Parameter vectors (pop_size, param_dim)
            
        Returns:
            fitness_scores: Fitness for each individual (pop_size,)
        """
        pop_size = params_batch.shape[0]
        
        # Evaluate each individual using vectorized operations
        fitness_scores = jax.vmap(
            lambda params: self._evaluate_individual(policy, params)
        )(params_batch)
        
        return fitness_scores

Verbatim archive excerpt from evojax_tests.ipynb. Context-dependent historical code, not a standalone runnable program. Comments retain their original wording; the article distinguishes implemented behavior from stale or overbroad comments.

The boundary that matters

The notebook’s own expansion of the PGPE acronym is not authoritative; the article uses the established algorithm name without repeating it. Device detection also does not prove that every operation runs efficiently on that backend. Warm-up, compilation, host copies and candidate memory belong in a performance measurement.

Keep building

Other posts of interest