[Pather] improve cost calculation for S-first

This commit is contained in:
Jan Petykiewicz 2026-07-16 16:28:56 -07:00
commit 1da5ac550a
4 changed files with 174 additions and 25 deletions

View file

@ -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.

View file

@ -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]] = []
solved = self.solve_parameters(steps, solve_indices, route_constraints) errors: list[Exception] = []
for solve_order, solve_indices in enumerate(solve_sets):
try:
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
order = self.order cost = sum(step.cost for step in selected_steps)
self.order += 1 feasible.append((float(cost), solve_order, tuple(selected_steps), end_port))
public_length = float(end_port.x) if self.request.length is None else float(self.request.length)
return Candidate( if not feasible:
tuple(selected_steps), if errors:
end_port, raise errors[-1]
sum(step.cost for step in selected_steps), raise BuildError(f'{self.request.route_name} composed primitive route is unsupported')
order,
public_length, minimum_cost = min(result[0] for result in feasible)
) cost_tied = [result for result in feasible if costs_equal(result[0], minimum_cost)]
raise BuildError(f'{self.request.route_name} composed primitive route is unsupported') cost, _solve_order, selected_steps, end_port = min(cost_tied, key=lambda result: result[1])
order = self.order
self.order += 1
public_length = float(end_port.x) if self.request.length is None else float(self.request.length)
return Candidate(selected_steps, end_port, cost, order, public_length)
@dataclass(frozen=True, slots=True) @dataclass(frozen=True, slots=True)

View file

@ -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)

View file

@ -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