From c7976b137e31492867da216b2a08dbbc0b53c92f Mon Sep 17 00:00:00 2001 From: Remco de Boer <29308176+redeboer@users.noreply.github.com> Date: Sun, 12 Jul 2026 15:43:42 +0200 Subject: [PATCH] BREAK: move particle-database matching out of `CSPSolver` --- src/qrules/solving.py | 165 +++++++++++++++++++++++-------------- src/qrules/workflow.py | 11 ++- tests/unit/test_solving.py | 21 +++-- 3 files changed, 125 insertions(+), 72 deletions(-) diff --git a/src/qrules/solving.py b/src/qrules/solving.py index f878f84d..b43499e7 100644 --- a/src/qrules/solving.py +++ b/src/qrules/solving.py @@ -227,19 +227,20 @@ def filter_quantum_number_problem_set( QuantumNumberSolution = MutableTransition[GraphEdgePropertyMap, GraphNodePropertyMap] +def _get_rule_name(rule: Any) -> str: + if inspect.isfunction(rule): + return rule.__name__ + if isinstance(rule, str): + return rule + return type(rule).__name__ + + def _convert_violated_rules_to_names( rules: dict[int, set[Rule]] | dict[int, set[GraphElementRule]], ) -> dict[int, set[str]]: - def get_name(rule: Any) -> str: - if inspect.isfunction(rule): - return rule.__name__ - if isinstance(rule, str): - return rule - return type(rule).__name__ - converted_dict = defaultdict(set) for node_id, rule_set in rules.items(): - converted_dict[node_id] = {get_name(rule) for rule in rule_set} + converted_dict[node_id] = {_get_rule_name(rule) for rule in rule_set} return converted_dict @@ -247,15 +248,8 @@ def get_name(rule: Any) -> str: def _convert_non_executed_rules_to_names( rules: dict[int, set[Rule]] | dict[int, set[GraphElementRule]], ) -> dict[int, set[str]]: - def get_name(rule: Any) -> str: - if inspect.isfunction(rule): - return rule.__name__ - if isinstance(rule, str): - return rule - return type(rule).__name__ - return { - node_id: {get_name(rule) for rule in rule_set} + node_id: {_get_rule_name(rule) for rule in rule_set} for node_id, rule_set in rules.items() } @@ -328,6 +322,86 @@ def find_solutions(self, problem_set: QNProblemSet) -> QNResult: """ +def complete_intermediate_states( + qn_result: QNResult, + qn_problem_set: QNProblemSet, + allowed_intermediate_states: Iterable[GraphEdgePropertyMap], +) -> QNResult: + """Match the intermediate states of a `QNResult` against a set of allowed states. + + A `Solver` assigns to each intermediate edge only the quantum numbers that were + solved for. This post-processing step completes the solutions: each intermediate + state is replaced by the full quantum number property maps of the matching entries + in :code:`allowed_intermediate_states` (one new solution per match, as e.g. created + by `.create_edge_properties` from a `.ParticleCollection`) and solutions without + any match are discarded. Conservation rules that the solver could not execute are + re-validated against the completed states, which may provide the previously missing + quantum numbers. + """ + topology = qn_problem_set.topology + completed_solutions = _insert_allowed_states( + qn_result.solutions, topology, allowed_intermediate_states + ) + node_not_executed_rules = _find_rules_by_name( + qn_problem_set.solving_settings.interactions, + qn_result.not_executed_node_rules, + ) + edge_not_executed_rules = _find_rules_by_name( + qn_problem_set.solving_settings.states, + qn_result.not_executed_edge_rules, + ) + if completed_solutions and (node_not_executed_rules or edge_not_executed_rules): + # rerun solver on these graphs using not executed rules and combine results + result = QNResult() + for completed_solution in completed_solutions: + interactions = completed_solution.interactions + states = completed_solution.states + interactions.update(qn_problem_set.initial_facts.interactions) + states.update(qn_problem_set.initial_facts.states) + result.extend( + validate_full_solution( + QNProblemSet( + initial_facts=MutableTransition(topology, states, interactions), + solving_settings=MutableTransition( + topology, + interactions={ + i: NodeSettings(conservation_rules=rules) # type: ignore[misc] + for i, rules in node_not_executed_rules.items() + }, + states={ + i: EdgeSettings(conservation_rules=rules) # type: ignore[misc] + for i, rules in edge_not_executed_rules.items() + }, + ), + ) + ) + ) + return result + return QNResult( + completed_solutions, + qn_result.not_executed_node_rules, + qn_result.violated_node_rules, + qn_result.not_executed_edge_rules, + qn_result.violated_edge_rules, + ) + + +def _find_rules_by_name( + settings: dict[int, NodeSettings] | dict[int, EdgeSettings], + rule_names: dict[int, set[str]], +) -> dict[int, set[Rule]]: + """Look up rule objects in graph element settings by their reported name.""" + return { + element_id: { + rule + for rule in settings[element_id].conservation_rules + if _get_rule_name(rule) in names + } + for element_id, names in rule_names.items() + if element_id in settings + } + + def _insert_allowed_states( solutions: list[QuantumNumberSolution], topology: Topology, @@ -338,7 +412,9 @@ def _insert_allowed_states( for solution in solutions: current_substituted_graphs = [solution] for edge_id in topology.intermediate_edge_ids: - incomplete_state = solution.states[edge_id] + incomplete_state = solution.states.get(edge_id) + if incomplete_state is None: + continue candidate_states = __get_candidate_states(incomplete_state, allowed_states) if len(candidate_states) == 0: message = f"Did not find any QN state candidate for edge id: {edge_id}" @@ -562,11 +638,13 @@ class CSPSolver(Solver): quantum numbers which are attributed to the interaction nodes (such as angular momentum :math:`L`). The conservation rules serve as the constraints and a special wrapper class serves as an adapter. + + The solutions carry only the quantum numbers that were solved for. Use + `complete_intermediate_states` to match the intermediate states against a set of + allowed states, such as the entries of a particle database. """ - def __init__( - self, allowed_intermediate_states: Iterable[GraphEdgePropertyMap] - ) -> None: + def __init__(self) -> None: self.__variables: set[_EdgeVariableInfo | _NodeVariableInfo] = set() self.__var_string_to_data: dict[str, _EdgeVariableInfo | _NodeVariableInfo] = {} self.__node_rules: dict[int, set[Rule]] = defaultdict(set) @@ -576,10 +654,9 @@ def __init__( defaultdict(set) ) self.__problem = Problem(BacktrackingSolver(forwardcheck=True)) - self.__allowed_intermediate_states = tuple(allowed_intermediate_states) self.__scoresheet = Scoresheet() - def find_solutions(self, problem_set: QNProblemSet) -> QNResult: # ruff: ignore[complex-structure] + def find_solutions(self, problem_set: QNProblemSet) -> QNResult: self.__initialize_constraints(problem_set) solutions = self.__problem.getSolutions() @@ -603,15 +680,8 @@ def find_solutions(self, problem_set: QNProblemSet) -> QNResult: # ruff: ignore solutions = self.__convert_solution_keys(problem_set.topology, solutions) - # insert particle instances - if self.__node_rules or self.__edge_rules: - selected_solutions = _insert_allowed_states( - solutions, - problem_set.topology, - self.__allowed_intermediate_states, - ) - else: - selected_solutions = [ + if not self.__node_rules and not self.__edge_rules: + solutions = [ QuantumNumberSolution( topology=problem_set.topology, interactions=problem_set.initial_facts.interactions, @@ -619,39 +689,8 @@ def find_solutions(self, problem_set: QNProblemSet) -> QNResult: # ruff: ignore ) ] - if selected_solutions and (node_not_executed_rules or edge_not_executed_rules): - # rerun solver on these graphs using not executed rules and combine results - topology = problem_set.topology - result = QNResult() - for full_particle_solution in selected_solutions: - interactions = full_particle_solution.interactions - states = full_particle_solution.states - interactions.update(problem_set.initial_facts.interactions) - states.update(problem_set.initial_facts.states) - result.extend( - validate_full_solution( - QNProblemSet( - initial_facts=MutableTransition( - topology, states, interactions - ), - solving_settings=MutableTransition( - topology, - interactions={ - i: NodeSettings(conservation_rules=rules) - for i, rules in node_not_executed_rules.items() - }, - states={ - i: EdgeSettings(conservation_rules=rules) - for i, rules in edge_not_executed_rules.items() - }, - ), - ) - ) - ) - return result - return QNResult( - selected_solutions, + solutions, _convert_non_executed_rules_to_names(node_not_executed_rules), _convert_violated_rules_to_names(node_not_satisfied_rules), _convert_non_executed_rules_to_names(edge_not_executed_rules), diff --git a/src/qrules/workflow.py b/src/qrules/workflow.py index afa07522..594b6eba 100644 --- a/src/qrules/workflow.py +++ b/src/qrules/workflow.py @@ -49,7 +49,7 @@ NumberOfThreads, create_interaction_settings, ) -from qrules.solving import CSPSolver +from qrules.solving import CSPSolver, complete_intermediate_states from qrules.system_control import ( GammaCheck, InteractionDeterminator, @@ -429,9 +429,12 @@ def _solve_single_problem( qn_problem_set: QNProblemSet, allowed_intermediate_states: Iterable[GraphEdgePropertyMap], ) -> tuple[QNProblemSet, QNResult]: - solver = CSPSolver(allowed_intermediate_states) - solutions = solver.find_solutions(qn_problem_set) - return qn_problem_set, solutions + solver = CSPSolver() + qn_solutions = solver.find_solutions(qn_problem_set) + completed_solutions = complete_intermediate_states( + qn_solutions, qn_problem_set, allowed_intermediate_states + ) + return qn_problem_set, completed_solutions def solve( diff --git a/tests/unit/test_solving.py b/tests/unit/test_solving.py index 655ac267..8646f6a2 100644 --- a/tests/unit/test_solving.py +++ b/tests/unit/test_solving.py @@ -13,7 +13,12 @@ spin_validity, ) from qrules.quantum_numbers import EdgeQuantumNumbers, NodeQuantumNumbers -from qrules.solving import CSPSolver, QNProblemSet, filter_quantum_number_problem_set +from qrules.solving import ( + CSPSolver, + QNProblemSet, + complete_intermediate_states, + filter_quantum_number_problem_set, +) if TYPE_CHECKING: from qrules.argument_handling import GraphEdgePropertyMap @@ -24,8 +29,11 @@ def it_finds_solutions( all_particles: qrules.particle.ParticleCollection, quantum_number_problem_set: QNProblemSet, ) -> None: - solver = CSPSolver(all_particles) - result = solver.find_solutions(quantum_number_problem_set) + solver = CSPSolver() + qn_result = solver.find_solutions(quantum_number_problem_set) + result = complete_intermediate_states( + qn_result, quantum_number_problem_set, all_particles + ) assert len(result.solutions) == 19 @pytest.mark.parametrize("with_spin_projection", [True, False]) @@ -34,7 +42,7 @@ def it_with_filtered_quantum_number_problem_set( quantum_number_problem_set: QNProblemSet, with_spin_projection: bool, ) -> None: - solver = CSPSolver(all_particles) + solver = CSPSolver() parametrized_edge_properties_and_domains = { EdgeQuantumNumbers.pid, # had to be added for c_parity_conservation to work EdgeQuantumNumbers.spin_magnitude, @@ -60,7 +68,10 @@ def it_with_filtered_quantum_number_problem_set( NodeQuantumNumbers.s_magnitude, ), ) - result = solver.find_solutions(new_quantum_number_problem_set) + qn_result = solver.find_solutions(new_quantum_number_problem_set) + result = complete_intermediate_states( + qn_result, new_quantum_number_problem_set, all_particles + ) if with_spin_projection: assert len(result.solutions) == 319