[Pather] improve cost calculation for S-first
This commit is contained in:
parent
ac34108253
commit
1da5ac550a
4 changed files with 174 additions and 25 deletions
|
|
@ -1127,8 +1127,9 @@ class Pather(PortList):
|
||||||
three-bend routes; other families try zero-to-two-bend routes before
|
three-bend routes; other families try zero-to-two-bend routes before
|
||||||
four-bend routes. The first band with a legal candidate wins. Within a
|
four-bend routes. The first band with a legal candidate wins. Within a
|
||||||
band, candidates are ordered by total cost, adapter count, step count,
|
band, candidates are ordered by total cost, adapter count, step count,
|
||||||
and deterministic discovery order. `strategy` controls straight-first
|
the requested straight-vs-turn topology preference, and deterministic
|
||||||
vs turn-first discovery order and therefore only the final tie-break.
|
discovery order. `strategy` therefore affects only otherwise tied
|
||||||
|
candidates.
|
||||||
Custom planning options may be supplied through `tool_options`; they
|
Custom planning options may be supplied through `tool_options`; they
|
||||||
are forwarded only to primitive offer generation.
|
are forwarded only to primitive offer generation.
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -32,6 +32,7 @@ from __future__ import annotations
|
||||||
from collections.abc import Iterable, Mapping, Sequence
|
from collections.abc import Iterable, Mapping, Sequence
|
||||||
from dataclasses import dataclass, replace
|
from dataclasses import dataclass, replace
|
||||||
from itertools import combinations
|
from itertools import combinations
|
||||||
|
from math import isclose as math_isclose
|
||||||
from typing import Any, Literal
|
from typing import Any, Literal
|
||||||
|
|
||||||
import numpy
|
import numpy
|
||||||
|
|
@ -71,6 +72,8 @@ from .interface import (
|
||||||
)
|
)
|
||||||
|
|
||||||
RouteTieBreakStrategy = Literal['straight_first', 'turn_first']
|
RouteTieBreakStrategy = Literal['straight_first', 'turn_first']
|
||||||
|
COST_RTOL = 1e-10
|
||||||
|
COST_ATOL = 1e-8
|
||||||
|
|
||||||
|
|
||||||
class NoLegalRouteError(BuildError):
|
class NoLegalRouteError(BuildError):
|
||||||
|
|
@ -89,6 +92,11 @@ def is_close(a: float, b: float) -> bool:
|
||||||
return scalar_close(a, b)
|
return scalar_close(a, b)
|
||||||
|
|
||||||
|
|
||||||
|
def costs_equal(a: float, b: float) -> bool:
|
||||||
|
"""Treat sub-resolution solver noise as equal during cost ranking."""
|
||||||
|
return math_isclose(float(a), float(b), rel_tol=COST_RTOL, abs_tol=COST_ATOL)
|
||||||
|
|
||||||
|
|
||||||
def clean_parameter(value: float) -> float:
|
def clean_parameter(value: float) -> float:
|
||||||
"""Snap tiny solver noise in primitive parameters before domain checks."""
|
"""Snap tiny solver noise in primitive parameters before domain checks."""
|
||||||
rounded = round(float(value))
|
rounded = round(float(value))
|
||||||
|
|
@ -274,7 +282,7 @@ class SolverRequest:
|
||||||
max_bends: int | None = None
|
max_bends: int | None = None
|
||||||
"""Optional override for grammar bend budget."""
|
"""Optional override for grammar bend budget."""
|
||||||
strategy: RouteTieBreakStrategy = 'straight_first'
|
strategy: RouteTieBreakStrategy = 'straight_first'
|
||||||
"""Discovery-order tie-break strategy for straight-vs-turn placement."""
|
"""Final topology tie-break preference for straight-vs-turn placement."""
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def route_name(self) -> str:
|
def route_name(self) -> str:
|
||||||
|
|
@ -339,6 +347,15 @@ class Solver:
|
||||||
count += 2
|
count += 2
|
||||||
return count
|
return count
|
||||||
|
|
||||||
|
def strategy_rank(self, steps: Sequence[SelectedPrimitive]) -> tuple[int, ...]:
|
||||||
|
"""Lexicographically rank every main straight-vs-turn placement."""
|
||||||
|
preferred_kind = 'straight' if self.request.strategy == 'straight_first' else 'turn'
|
||||||
|
return tuple(
|
||||||
|
int(('straight' if step.offer.kind == 'straight' else 'turn') != preferred_kind)
|
||||||
|
for step in steps
|
||||||
|
if step.role != 'adapter'
|
||||||
|
)
|
||||||
|
|
||||||
def candidate_key(self, candidate: Candidate) -> tuple[Any, ...]:
|
def candidate_key(self, candidate: Candidate) -> tuple[Any, ...]:
|
||||||
"""Return a deterministic key for duplicate solved candidates."""
|
"""Return a deterministic key for duplicate solved candidates."""
|
||||||
def endpoint_key(port: Port) -> tuple[float, float, float | None, str | None]:
|
def endpoint_key(port: Port) -> tuple[float, float, float | None, str | None]:
|
||||||
|
|
@ -425,12 +442,14 @@ class Solver:
|
||||||
) from last_error
|
) from last_error
|
||||||
raise NoLegalRouteError(f'No legal primitive offer for {self.request.route_name}')
|
raise NoLegalRouteError(f'No legal primitive offer for {self.request.route_name}')
|
||||||
|
|
||||||
|
minimum_cost = min(candidate.cost for candidate in candidates)
|
||||||
|
cost_tied = [candidate for candidate in candidates if costs_equal(candidate.cost, minimum_cost)]
|
||||||
return min(
|
return min(
|
||||||
candidates,
|
cost_tied,
|
||||||
key=lambda candidate: (
|
key=lambda candidate: (
|
||||||
round(float(candidate.cost), 9),
|
|
||||||
sum(step.role == 'adapter' for step in candidate.steps),
|
sum(step.role == 'adapter' for step in candidate.steps),
|
||||||
len(candidate.steps),
|
len(candidate.steps),
|
||||||
|
self.strategy_rank(candidate.steps),
|
||||||
candidate.order,
|
candidate.order,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
@ -793,10 +812,7 @@ class Solver:
|
||||||
residual_jog: float,
|
residual_jog: float,
|
||||||
) -> Iterable[tuple[SelectedPrimitive, ...]]:
|
) -> Iterable[tuple[SelectedPrimitive, ...]]:
|
||||||
"""Recursively enumerate normal/adapter/turn blocks within the bend budget."""
|
"""Recursively enumerate normal/adapter/turn blocks within the bend budget."""
|
||||||
straight_options = self.straight_options(steps)
|
for normal in self.straight_options(steps):
|
||||||
if self.request.strategy == 'straight_first':
|
|
||||||
straight_options = (*straight_options[1:], straight_options[0])
|
|
||||||
for normal in straight_options:
|
|
||||||
after_normal = (*steps, *normal)
|
after_normal = (*steps, *normal)
|
||||||
suffix_options = (
|
suffix_options = (
|
||||||
self.adapter_options(after_normal, residual_jog=0)
|
self.adapter_options(after_normal, residual_jog=0)
|
||||||
|
|
@ -955,10 +971,10 @@ class Solver:
|
||||||
|
|
||||||
def finalize(self, steps: Sequence[SelectedPrimitive]) -> Candidate:
|
def finalize(self, steps: Sequence[SelectedPrimitive]) -> Candidate:
|
||||||
"""
|
"""
|
||||||
Try all small solve sets for one raw sequence and return the first match.
|
Try all small solve sets for one raw sequence and return the cheapest match.
|
||||||
|
|
||||||
Solve-set order is deterministic and becomes part of the candidate
|
Recoverable failures reject only the current solve set. Solve-set order
|
||||||
ordering only after cost and structural tie-breakers.
|
breaks ties between equally priced parameter allocations.
|
||||||
"""
|
"""
|
||||||
constraints: list[tuple[Literal['x', 'y'], float]] = []
|
constraints: list[tuple[Literal['x', 'y'], float]] = []
|
||||||
if self.request.length is not None:
|
if self.request.length is not None:
|
||||||
|
|
@ -978,24 +994,35 @@ class Solver:
|
||||||
for solve_size in range(1, max_solve + 1):
|
for solve_size in range(1, max_solve + 1):
|
||||||
solve_sets.extend(combinations(adjustable, solve_size))
|
solve_sets.extend(combinations(adjustable, solve_size))
|
||||||
|
|
||||||
for solve_indices in solve_sets:
|
feasible: list[tuple[float, int, tuple[SelectedPrimitive, ...], Port]] = []
|
||||||
|
errors: list[Exception] = []
|
||||||
|
for solve_order, solve_indices in enumerate(solve_sets):
|
||||||
|
try:
|
||||||
solved = self.solve_parameters(steps, solve_indices, route_constraints)
|
solved = self.solve_parameters(steps, solve_indices, route_constraints)
|
||||||
|
except (BuildError, NotImplementedError, PortError) as err:
|
||||||
|
raise_if_fatal(err)
|
||||||
|
errors.append(err)
|
||||||
|
continue
|
||||||
if solved is None:
|
if solved is None:
|
||||||
continue
|
continue
|
||||||
selected_steps, end_port = solved
|
selected_steps, end_port = solved
|
||||||
if not self.endpoint_matches(end_port, route_constraints):
|
if not self.endpoint_matches(end_port, route_constraints):
|
||||||
continue
|
continue
|
||||||
|
cost = sum(step.cost for step in selected_steps)
|
||||||
|
feasible.append((float(cost), solve_order, tuple(selected_steps), end_port))
|
||||||
|
|
||||||
|
if not feasible:
|
||||||
|
if errors:
|
||||||
|
raise errors[-1]
|
||||||
|
raise BuildError(f'{self.request.route_name} composed primitive route is unsupported')
|
||||||
|
|
||||||
|
minimum_cost = min(result[0] for result in feasible)
|
||||||
|
cost_tied = [result for result in feasible if costs_equal(result[0], minimum_cost)]
|
||||||
|
cost, _solve_order, selected_steps, end_port = min(cost_tied, key=lambda result: result[1])
|
||||||
order = self.order
|
order = self.order
|
||||||
self.order += 1
|
self.order += 1
|
||||||
public_length = float(end_port.x) if self.request.length is None else float(self.request.length)
|
public_length = float(end_port.x) if self.request.length is None else float(self.request.length)
|
||||||
return Candidate(
|
return Candidate(selected_steps, end_port, cost, order, public_length)
|
||||||
tuple(selected_steps),
|
|
||||||
end_port,
|
|
||||||
sum(step.cost for step in selected_steps),
|
|
||||||
order,
|
|
||||||
public_length,
|
|
||||||
)
|
|
||||||
raise BuildError(f'{self.request.route_name} composed primitive route is unsupported')
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True, slots=True)
|
@dataclass(frozen=True, slots=True)
|
||||||
|
|
|
||||||
|
|
@ -950,6 +950,59 @@ def test_autotool_sbend_explicit_cost_can_override_geometric_cost() -> None:
|
||||||
assert_allclose(out_port.offset, [20, 4])
|
assert_allclose(out_port.offset, [20, 4])
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize('jog', [41_988.0, -41_988.0])
|
||||||
|
def test_autotool_strategy_orders_main_steps_across_adapters(jog: float) -> None:
|
||||||
|
radius = 60_000.0
|
||||||
|
transition_length = 50_000.0
|
||||||
|
route_length = 545_929.0
|
||||||
|
primary = 'primary'
|
||||||
|
secondary = 'secondary'
|
||||||
|
sbend_endpoint = circular_arc_sbend_endpoint(radius, primary)
|
||||||
|
|
||||||
|
def make_primary_sbend(offset: float) -> Pattern:
|
||||||
|
pattern = Pattern()
|
||||||
|
pattern.ports['A'] = Port((0, 0), 0, ptype=primary)
|
||||||
|
pattern.ports['B'] = sbend_endpoint(offset)
|
||||||
|
return pattern
|
||||||
|
|
||||||
|
library = Library()
|
||||||
|
library['bend'] = make_bend(radius, ptype=primary)
|
||||||
|
transition = Pattern()
|
||||||
|
transition.ports['P'] = Port((0, 0), 0, ptype=primary)
|
||||||
|
transition.ports['S'] = Port((transition_length, 0), pi, ptype=secondary)
|
||||||
|
library['transition'] = transition
|
||||||
|
tool = (
|
||||||
|
AutoTool(bbox_library=library)
|
||||||
|
.add_straight(lambda length: make_straight(length, ptype=primary), primary, 'A')
|
||||||
|
.add_straight(
|
||||||
|
lambda length: make_straight(length, ptype=secondary),
|
||||||
|
secondary,
|
||||||
|
'A',
|
||||||
|
cost=0.3,
|
||||||
|
)
|
||||||
|
.add_bend(library.abstract('bend'), 'A', 'B', clockwise=True)
|
||||||
|
.add_sbend(
|
||||||
|
make_primary_sbend,
|
||||||
|
primary,
|
||||||
|
'A',
|
||||||
|
'B',
|
||||||
|
jog_range=(0, 2 * radius),
|
||||||
|
endpoint=sbend_endpoint,
|
||||||
|
)
|
||||||
|
.add_transition(library.abstract('transition'), 'P', 'S')
|
||||||
|
)
|
||||||
|
|
||||||
|
selected_kinds = {}
|
||||||
|
for strategy in ('straight_first', 'turn_first'):
|
||||||
|
pather = Pather(library, tools=tool, render='deferred')
|
||||||
|
pather.ports['A'] = Port((0, 0), 0, ptype=primary)
|
||||||
|
pather.jog('A', jog, length=route_length, out_ptype=primary, strategy=strategy)
|
||||||
|
selected_kinds[strategy] = [step.kind for step in pather._paths['A']]
|
||||||
|
|
||||||
|
assert selected_kinds['straight_first'] == ['straight', 'straight', 'straight', 's']
|
||||||
|
assert selected_kinds['turn_first'] == ['s', 'straight', 'straight', 'straight']
|
||||||
|
|
||||||
|
|
||||||
def test_autotool_add_methods_propagate_callable_cost_to_all_created_offers() -> None:
|
def test_autotool_add_methods_propagate_callable_cost_to_all_created_offers() -> None:
|
||||||
def cost(parameter: float, endpoint: Port) -> float:
|
def cost(parameter: float, endpoint: Port) -> float:
|
||||||
return abs(parameter) + abs(endpoint.x) + abs(endpoint.y)
|
return abs(parameter) + abs(endpoint.x) + abs(endpoint.y)
|
||||||
|
|
|
||||||
|
|
@ -122,6 +122,74 @@ def test_solver_offer_cache_accepts_unhashable_request_tool_options() -> None:
|
||||||
assert tool.calls == 1
|
assert tool.calls == 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_solver_finalize_chooses_cheapest_parameter_allocation() -> None:
|
||||||
|
tool = PlanningOnlyTool()
|
||||||
|
solver = Solver(SolverRequest(
|
||||||
|
family='straight',
|
||||||
|
tool=tool,
|
||||||
|
in_ptype='wire',
|
||||||
|
tool_options={},
|
||||||
|
length=10,
|
||||||
|
))
|
||||||
|
expensive_offer = StraightOffer.generated('wire', lambda length: length, cost=1)
|
||||||
|
cheap_offer = StraightOffer.generated('wire', lambda length: length, cost=0.3)
|
||||||
|
expensive = solver.evaluate(expensive_offer, 0, 'wire', out_ptype=None, role='main')
|
||||||
|
cheap = solver.evaluate(cheap_offer, 0, 'wire', out_ptype=None, role='main')
|
||||||
|
|
||||||
|
for steps in ((expensive, cheap), (cheap, expensive)):
|
||||||
|
candidate = solver.finalize(steps)
|
||||||
|
parameters = {step.offer: step.parameter for step in candidate.steps}
|
||||||
|
|
||||||
|
assert parameters[cheap_offer] == pytest.approx(10)
|
||||||
|
assert parameters[expensive_offer] == pytest.approx(0)
|
||||||
|
assert candidate.cost == pytest.approx(3)
|
||||||
|
|
||||||
|
|
||||||
|
def test_solver_strategy_rank_covers_all_main_steps_and_ignores_adapters() -> None:
|
||||||
|
tool = PlanningOnlyTool()
|
||||||
|
straight_first = Solver(SolverRequest(
|
||||||
|
family='s',
|
||||||
|
tool=tool,
|
||||||
|
in_ptype='wire',
|
||||||
|
tool_options={},
|
||||||
|
strategy='straight_first',
|
||||||
|
))
|
||||||
|
turn_first = Solver(SolverRequest(
|
||||||
|
family='s',
|
||||||
|
tool=tool,
|
||||||
|
in_ptype='wire',
|
||||||
|
tool_options={},
|
||||||
|
strategy='turn_first',
|
||||||
|
))
|
||||||
|
straight_offer = StraightOffer.generated('wire', lambda length: length)
|
||||||
|
bend_offer = BendOffer.prebuilt(
|
||||||
|
'wire',
|
||||||
|
'wire',
|
||||||
|
Port((1, 1), 3 * pi / 2, ptype='wire'),
|
||||||
|
None,
|
||||||
|
ccw=True,
|
||||||
|
)
|
||||||
|
adapter_offer = StraightOffer.prebuilt(
|
||||||
|
'wire',
|
||||||
|
'adapted',
|
||||||
|
Port((1, 0), pi, ptype='adapted'),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
straight = straight_first.evaluate(straight_offer, 0, 'wire', out_ptype=None, role='main')
|
||||||
|
turn = straight_first.evaluate(bend_offer, 1, 'wire', out_ptype=None, role='main')
|
||||||
|
adapter = straight_first.evaluate(adapter_offer, 1, 'wire', out_ptype=None, role='adapter')
|
||||||
|
alternating = (straight, turn, adapter, straight, turn)
|
||||||
|
delayed_straight = (straight, turn, adapter, turn, straight)
|
||||||
|
|
||||||
|
assert straight_first.strategy_rank(alternating) < straight_first.strategy_rank(delayed_straight)
|
||||||
|
assert turn_first.strategy_rank(delayed_straight) < turn_first.strategy_rank(alternating)
|
||||||
|
assert straight_first.strategy_rank(alternating) == straight_first.strategy_rank(
|
||||||
|
(straight, turn, straight, turn),
|
||||||
|
)
|
||||||
|
assert straight_first.strategy_rank((straight,)) < straight_first.strategy_rank((turn,))
|
||||||
|
assert turn_first.strategy_rank((turn,)) < turn_first.strategy_rank((straight,))
|
||||||
|
|
||||||
|
|
||||||
def test_tool_requires_primitive_offers_override() -> None:
|
def test_tool_requires_primitive_offers_override() -> None:
|
||||||
class RenderOnlyTool(Tool):
|
class RenderOnlyTool(Tool):
|
||||||
def render(self, batch, *, port_names=('A', 'B'), **kwargs) -> Library: # noqa: ANN001,ANN202,ARG002
|
def render(self, batch, *, port_names=('A', 'B'), **kwargs) -> Library: # noqa: ANN001,ANN202,ARG002
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue