[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

@ -950,6 +950,59 @@ def test_autotool_sbend_explicit_cost_can_override_geometric_cost() -> None:
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 cost(parameter: float, endpoint: Port) -> float:
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
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:
class RenderOnlyTool(Tool):
def render(self, batch, *, port_names=('A', 'B'), **kwargs) -> Library: # noqa: ANN001,ANN202,ARG002