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
10 changes: 5 additions & 5 deletions vizier/_src/algorithms/evolution/nsga2.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,7 @@
Mutation = templates.Mutation


def _pareto_rank(ys: np.ndarray) -> np.ndarray:
def pareto_rank(ys: np.ndarray) -> np.ndarray:
"""Pareto rank, which is the number of points dominating it.

Args:
Expand All @@ -44,7 +44,7 @@ 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) -> np.ndarray:
"""Crowding distance.

Args:
Expand Down Expand Up @@ -123,7 +123,7 @@ def __init__(
self,
target_size: int,
*,
ranking_fn: Callable[[np.ndarray], np.ndarray] = _pareto_rank,
ranking_fn: Callable[[np.ndarray], np.ndarray] = pareto_rank,
eviction_limit: Optional[int] = None
):
"""Init.
Expand Down Expand Up @@ -183,7 +183,7 @@ def select(self, population: Population) -> 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)
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 All @@ -207,7 +207,7 @@ def __init__(
population_size: int = 50,
first_survival_after: Optional[int] = None,
*,
ranking_fn: Callable[[np.ndarray], np.ndarray] = _pareto_rank,
ranking_fn: Callable[[np.ndarray], np.ndarray] = pareto_rank,
eviction_limit: Optional[int] = None,
adaptation: Optional[Mutation[Population, Offspring]] = None,
adaptation_callable: Optional[
Expand Down
14 changes: 7 additions & 7 deletions vizier/_src/algorithms/evolution/nsga2_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -69,30 +69,30 @@ class Nsga2Test(absltest.TestCase):

def test_pareto_rank_empty(self):
ys = np.array([]).reshape(0, 2)
ranks = nsga2._pareto_rank(ys)
ranks = nsga2.pareto_rank(ys)
self.assertEqual(ranks.shape, (0,))

def test_pareto_rank_single(self):
ys = np.array([[1.0, 2.0]])
ranks = nsga2._pareto_rank(ys)
ranks = nsga2.pareto_rank(ys)
np.testing.assert_array_equal(ranks, [0])

def test_pareto_rank_simple_dominance(self):
# P0 dominates P1
# P0 dominates P2
# P1 and P2 don't dominate each other
ys = np.array([[2.0, 3.0], [1.0, 3.0], [2.0, 2.0]])
ranks = nsga2._pareto_rank(ys)
ranks = nsga2.pareto_rank(ys)
np.testing.assert_array_equal(ranks, [0, 1, 1])

def test_pareto_rank_no_dominance(self):
ys = np.array([[1.0, 5.0], [2.0, 4.0], [3.0, 3.0]])
ranks = nsga2._pareto_rank(ys)
ranks = nsga2.pareto_rank(ys)
np.testing.assert_array_equal(ranks, [0, 0, 0])

def test_pareto_rank_duplicate_points_do_not_dominate_each_other(self):
ys = np.array([[2.0, 3.0], [1.0, 2.0], [2.0, 3.0]])
ranks = nsga2._pareto_rank(ys)
ranks = nsga2.pareto_rank(ys)
np.testing.assert_array_equal(ranks, [0, 2, 0])

def test_pareto_rank_larger_case(self):
Expand All @@ -103,7 +103,7 @@ def test_pareto_rank_larger_case(self):
[8, 4], # 3: (dominated by [10, 5], [8, 5], [9, 4])
[1, 10], # 0
])
ranks = nsga2._pareto_rank(ys)
ranks = nsga2.pareto_rank(ys)
np.testing.assert_array_equal(ranks, [0, 1, 1, 3, 0])

def test_survival_by_pareto_rank(self):
Expand Down Expand Up @@ -244,7 +244,7 @@ def test_comprehensive_sanity_check(self):
self.assertTrue(np.all(algorithm.population.ages <= 3))

ys = algorithm.population.ys
pareto = algorithm.population[nsga2._pareto_rank(ys) == 0]
pareto = algorithm.population[nsga2.pareto_rank(ys) == 0]
logging.info('Pareto frontier %s %s', pareto.xs, pareto.ys)

# Smoke test dump-load.
Expand Down
Loading