Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
70 changes: 59 additions & 11 deletions vizier/_src/algorithms/evolution/nsga2.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,8 +16,9 @@

"""NSGA-II algorithm: https://ieeexplore.ieee.org/document/996017."""

from typing import Callable, Optional, Tuple
from typing import Callable, Optional, Sequence, Tuple

from absl import logging
import attr
import numpy as np
from vizier import pyvizier as vz
Expand All @@ -44,30 +45,62 @@ def pareto_rank(ys: np.ndarray) -> np.ndarray:
return np.sum(np.stack(dominated), axis=0)


def crowding_distance(ys: np.ndarray) -> np.ndarray:
def crowding_distance(
ys: np.ndarray, *, extra_tiebreakers: Sequence[np.ndarray] = tuple()
) -> np.ndarray:
"""Crowding distance.

Reference:
https://medium.com/@rossleecooloh/optimization-algorithm-nsga-ii-and-python-package-deap-fca0be6b2ffc
except that the lower boundary does not get infinity crowding score.

Args:
ys: (number of population) x (number of metrics) array.
extra_tiebreakers: A sequence of (number of population) array of floating
numbers. If specified, they are used to break ties when sorting the
population. By default, a random number is always used to break ties
consistently across all metrics.

Returns:
(number of population) float32 array. Higher numbers mean less crowding
and more desirable.
"""
scores = np.zeros([ys.shape[0]], dtype=np.float32)

if ys.shape[0] <= 1:
return scores

rng = np.random.default_rng()
# Use a random number to break ties. But the same random number is used for
# all metrics.
random_tiebreaker = rng.random(ys.shape[0])
tiebreakers = list(extra_tiebreakers) + [random_tiebreaker]

for m in range(ys.shape[1]):
# Sort by the m-th metric.
sid = sorted(
np.arange(ys.shape[0]),
key=lambda i, m=m: (ys[i, m],)
+ tuple(tiebreaker[i] for tiebreaker in tiebreakers),
)

# Compute the range of the m-th metric.
yy = ys[:, m] # Shape: (num_population,)
sid = np.argsort(yy)
yrange = yy[sid[-1]] - yy[sid[0]] + np.finfo(np.float32).eps

# Boundary are assigned infinity.
scores[sid[0]] += np.inf
# Lower boundary is assigned a one-sided score and does not automatically
# get infinity. This is different from the paper. The lower boundary means
# it's dominated by all other points in one dimension. There's no reason to
# favor it over other points.
scores[sid[0]] += (yy[sid[1]] - yy[sid[0]]) / yrange
# Upper boundary is assigned infinity. This point will survive anyways
# because it's pareto-optimal. But in case there are ties, it's useful to
# make only one of them stand out.
scores[sid[-1]] += np.inf

# Compute the crowding distance.
yrange = yy[sid[-1]] - yy[sid[0]] + np.finfo(np.float32).eps
scores[sid[1:-1]] += (yy[sid[2:]] - yy[sid[:-2]]) / yrange
return scores
# Normalize the score to [0, 1].
return scores / ys.shape[1]


def _constraint_violation(ys: np.ndarray) -> np.ndarray:
Expand Down Expand Up @@ -163,6 +196,7 @@ def select(self, population: Population) -> Population:
# return empty.
return population

logging.info('Selecting %s from %s', self._target_size, len(population))
selected = population.empty_like()
# Sort by the safety constraint.
if selected.cs.shape[1]:
Expand All @@ -171,19 +205,33 @@ def select(self, population: Population) -> Population:
)
selected += population[top]
population = population[border]

logging.info(
'Selected %s by safety constraints, and will break ties among: %s',
len(selected),
len(population),
)
# Sort by the pareto rank.
pareto_ranks = self._ranking_fn(population.ys)
top, border = _select_by(
pareto_ranks, target=self._target_size - len(selected)
)
considered_pareto_ranks = np.concatenate(
[-np.ones(len(selected)), pareto_ranks[top], pareto_ranks[border]]
)
selected += population[top]
population = population[border]

logging.info(
'Selected %s by pareto rank, and will break ties among: %s',
len(selected),
len(population),
)
# Sort by the distance. Include the points that are already selected for
# the computation.
# Flip the sign so it works with ascending sort.
distance = -crowding_distance((selected + population).ys)
distance = -crowding_distance(
(selected + population).ys,
extra_tiebreakers=[considered_pareto_ranks],
)
sids = np.argsort(distance)
# Selected points have fewer constraint violations or better pareto rank.
# Regardless of the distance, they remain selected. Rank the remainder only.
Expand Down
28 changes: 25 additions & 3 deletions vizier/_src/algorithms/evolution/nsga2_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,27 @@
np.set_printoptions(precision=3)


class CrowdingDistanceTest(absltest.TestCase):

def test_crowding_distance(self):
ys = np.array([
[0.5, 4.0],
[0.5, 4.0],
[0.5, 2.0],
[1.0, 2.0],
[1.0, 2.0],
])
result = nsga2.crowding_distance(
ys, extra_tiebreakers=[nsga2.pareto_rank(ys)]
)
np.testing.assert_allclose(
# Only one of the duplicate pareto-optimal points should get
# infinity score.
np.isinf(result).sum(),
2,
)


def nsga2_on_all_types(
population_size: int = 50, eviction_limit: Optional[int] = None
) -> templates.CanonicalEvolutionDesigner[nsga2.Population, nsga2.Offspring]:
Expand Down Expand Up @@ -145,8 +166,9 @@ def test_survival_by_crowding_distance(self):
vz.Measurement({'m1': 1.001, 'm2': -1.001, 's1': 2.0, 's2': 0.0})
)

# 4 safe trials with the same pareto rank. Crowding distance is computed
# among them to break ties. Trial 3 is less "crowded" than Trial 2.
# 4 safe trials with the same pareto rank. Ties are broken by their crowding
# distances, computed *with* the pareto optimal trials. Trial 1 is most
# crowded because of its proximity to trial 0.
trial1 = vz.Trial(id=1)
trial1.complete(
vz.Measurement({'m1': 1.0, 'm2': 0.0, 's1': 2.0, 's2': 0.9})
Expand All @@ -166,7 +188,7 @@ def test_survival_by_crowding_distance(self):

trials = vza.CompletedTrials([trial0, trial1, trial2, trial3, trial4])
algorithm.update(trials, vza.ActiveTrials())
self.assertSetEqual(set(algorithm.population.trial_ids), {0, 1, 3, 4})
self.assertSetEqual(set(algorithm.population.trial_ids), {0, 2, 3, 4})

def test_survival_by_safety(self):
algorithm = nsga2_on_all_types(3)
Expand Down
Loading