diff --git a/src/qrules/io/_dict.py b/src/qrules/io/_dict.py index a3c92623..205a58a0 100644 --- a/src/qrules/io/_dict.py +++ b/src/qrules/io/_dict.py @@ -40,11 +40,18 @@ def from_attrs_decorated(inst: Any) -> dict: ) -def _value_serializer(inst: type, field: attrs.Attribute, value: Any) -> Any: # ruff: ignore[unused-function-argument, too-many-return-statements] +def _value_serializer(inst: type, field: attrs.Attribute, value: Any) -> Any: # ruff: ignore[too-many-return-statements] + if isinstance(value, (frozenset, set)): + return sorted(value) if isinstance(value, abc.Mapping): if all(isinstance(p, Particle) for p in value.values()): return {k: v.name for k, v in value.items()} - return dict(value) + return { + (key.__name__ if callable(key) else key): _value_serializer( + inst, field, item + ) + for key, item in value.items() + } if not isinstance(inst, (ReactionInfo, State, FrozenTransition)): # ruff: ignore[collapsible-if] if isinstance(value, Particle): return value.name diff --git a/src/qrules/solving.py b/src/qrules/solving.py index 99df2a7c..eab830d1 100644 --- a/src/qrules/solving.py +++ b/src/qrules/solving.py @@ -225,6 +225,68 @@ def filter_quantum_number_problem_set( ) +def strip_quantum_numbers( + problem_set: QNProblemSet, + edge_qns: Iterable[EdgeQuantumNumberTypes] = (), + node_qns: Iterable[NodeQuantumNumberTypes] = (), +) -> QNProblemSet: + """Remove the given quantum numbers from the facts and domains of a problem set. + + The conservation rules are kept: a rule that can no longer execute without the + removed quantum numbers is skipped and reported by the solver through the + not-executed-rules mechanism. + """ + edge_qn_set = set(edge_qns) + node_qn_set = set(node_qns) + facts = problem_set.initial_facts + settings = problem_set.solving_settings + new_facts = MutableTransition( + facts.topology, + states={ # type: ignore[arg-type] + edge_id: { + qn_type: value + for qn_type, value in prop_map.items() + if qn_type not in edge_qn_set + } + for edge_id, prop_map in facts.states.items() + }, + interactions={ # type: ignore[arg-type] + node_id: { + qn_type: value + for qn_type, value in prop_map.items() + if qn_type not in node_qn_set + } + for node_id, prop_map in facts.interactions.items() + }, + ) + new_settings = MutableTransition( + settings.topology, + states={ # type: ignore[arg-type] + edge_id: attrs.evolve( + edge_settings, + qn_domains={ + qn_type: domain + for qn_type, domain in edge_settings.qn_domains.items() + if qn_type not in edge_qn_set + }, + ) + for edge_id, edge_settings in settings.states.items() + }, + interactions={ # type: ignore[arg-type] + node_id: attrs.evolve( + node_settings, + qn_domains={ + qn_type: domain + for qn_type, domain in node_settings.qn_domains.items() + if qn_type not in node_qn_set + }, + ) + for node_id, node_settings in settings.interactions.items() + }, + ) + return QNProblemSet(initial_facts=new_facts, solving_settings=new_settings) + + def merge_qn_problem_sets( qn_problem_sets: Iterable[QNProblemSet], merge_qns: Iterable[EdgeQuantumNumberTypes] | None = None, diff --git a/src/qrules/workflow.py b/src/qrules/workflow.py index 7d56fbc5..95e0e057 100644 --- a/src/qrules/workflow.py +++ b/src/qrules/workflow.py @@ -25,7 +25,7 @@ from copy import copy, deepcopy from functools import partial from multiprocessing import Pool -from typing import TYPE_CHECKING, overload +from typing import TYPE_CHECKING, Any, overload from attrs import define, field, frozen from tqdm.auto import tqdm @@ -52,8 +52,10 @@ ) from qrules.solving import ( CSPSolver, + _create_merge_key, complete_intermediate_states, merge_qn_problem_sets, + strip_quantum_numbers, ) from qrules.system_control import ( GammaCheck, @@ -66,6 +68,8 @@ remove_duplicate_solutions, ) from qrules.topology import ( + FrozenDict, + FrozenTransition, MutableTransition, create_isobar_topologies, create_n_body_topology, @@ -848,6 +852,115 @@ def find_solutions( # ruff: ignore[too-many-positional-arguments] return collect_reaction_info(results, final_state, formalism) +QNTransition = FrozenTransition[FrozenDict[Any, Any], FrozenDict[Any, Any]] +"""Transition whose states and interactions are quantum-number property maps.""" + + +@overload +def strip_spin_projections( + qn_problem_sets: QNProblemSetCollection, +) -> QNProblemSetCollection: ... +@overload +def strip_spin_projections( + qn_problem_sets: dict[float, list[QNProblemSet]], +) -> dict[float, list[QNProblemSet]]: ... +def strip_spin_projections( + qn_problem_sets: QNProblemSetCollection | dict[float, list[QNProblemSet]], +) -> QNProblemSetCollection | dict[float, list[QNProblemSet]]: + """Remove spin projections from the problem sets, for :math:`J^{P(C)}`-level solving. + + Removes the spin projections of the initial and final state from the facts, the + spin-projection domains of the intermediate edges, and the :math:`L`/:math:`S` + projection domains of the interaction nodes. Problem sets that thereby become + identical — such as the expansion over all spin-projection combinations — are + deduplicated. Rules that require spin projections are skipped and reported by the + solver through the not-executed-rules mechanism. + """ + if isinstance(qn_problem_sets, QNProblemSetCollection): + stripped_collection = copy(qn_problem_sets) + stripped_collection.problem_sets = strip_spin_projections( + qn_problem_sets.problem_sets + ) + return stripped_collection + return { + strength: _unique_problem_sets( + strip_quantum_numbers( + problem_set, + edge_qns={EdgeQuantumNumbers.spin_projection}, + node_qns={ + NodeQuantumNumbers.l_projection, + NodeQuantumNumbers.s_projection, + }, + ) + for problem_set in problem_sets + ) + for strength, problem_sets in qn_problem_sets.items() + } + + +def _unique_problem_sets(problem_sets: Iterable[QNProblemSet]) -> list[QNProblemSet]: + unique: dict[tuple, QNProblemSet] = {} + for problem_set in problem_sets: + unique.setdefault(_create_merge_key(problem_set, set()), problem_set) + return list(unique.values()) + + +def find_qn_transitions( + qn_problem_sets: QNProblemSetCollection | dict[float, list[QNProblemSet]], +) -> tuple[QNTransition, ...]: + """Find allowed transitions purely at the quantum-number level. + + Solves the problem sets with the `.CSPSolver` without completing the intermediate + states from a particle database: the intermediate states of the returned + transitions carry exactly the quantum numbers that the problem sets declare as + domains. Combine with `strip_spin_projections` to obtain :math:`J^{P(C)}`-level + transitions for e.g. a Dalitz-plot decomposition. + """ + if isinstance(qn_problem_sets, QNProblemSetCollection): + qn_problem_sets = qn_problem_sets.problem_sets + solver = CSPSolver() + qn_results: dict[float, list[tuple[QNProblemSet, QNResult]]] = defaultdict(list) + for strength, qn_problems in sorted(qn_problem_sets.items(), reverse=True): + qn_results[strength].extend( + (qn_problem_set, solver.find_solutions(qn_problem_set)) + for qn_problem_set in qn_problems + ) + return collect_qn_transitions(qn_results) + + +def collect_qn_transitions( + qn_results: dict[float, list[tuple[QNProblemSet, QNResult]]], +) -> tuple[QNTransition, ...]: + """Summarize solver results as unique quantum-number-level transitions. + + The counterpart of `collect_reaction_info` for a quantum-number-only workflow: no + particle database is consulted and no `.State` objects are created — each + transition carries exactly the quantum numbers that are known from the initial + facts or were solved for. + """ + transitions: dict[QNTransition, None] = {} + for qn_result_pairs in qn_results.values(): + for qn_problem_set, qn_result in qn_result_pairs: + facts = qn_problem_set.initial_facts + for solution in qn_result.solutions: + states = dict(solution.states) + for edge_id, edge_facts in facts.states.items(): + states[edge_id] = {**edge_facts, **states.get(edge_id, {})} + interactions = dict(solution.interactions) + for node_id, node_facts in facts.interactions.items(): + interactions[node_id] = { + **node_facts, + **interactions.get(node_id, {}), + } + transition: QNTransition = FrozenTransition( + solution.topology, + states={i: FrozenDict(m) for i, m in states.items()}, + interactions={i: FrozenDict(m) for i, m in interactions.items()}, + ) + transitions[transition] = None + return tuple(transitions) + + def _match_final_state_ids( graph: MutableTransition[ParticleWithSpin, InteractionProperties], state_definition: Sequence[StateDefinition], diff --git a/tests/unit/test_workflow.py b/tests/unit/test_workflow.py index 15f5e6a8..d68aad78 100644 --- a/tests/unit/test_workflow.py +++ b/tests/unit/test_workflow.py @@ -1,5 +1,8 @@ +import json + import pytest +from qrules.io import asdict from qrules.particle import ParticleCollection, load_pdg from qrules.quantum_numbers import EdgeQuantumNumbers from qrules.settings import ( @@ -13,7 +16,9 @@ QNProblemSetCollection, create_qn_problem_sets, filter_intermediate_particles, + find_qn_transitions, find_solutions, + strip_spin_projections, ) @@ -188,3 +193,55 @@ def test_pipeline_reproduces_state_transition_manager( ) workflow_reaction = find_solutions(qn_problem_sets, particle_db) assert workflow_reaction == reaction + + +def test_projection_free_qn_transitions(): + particle_db = load_pdg() + collection = create_qn_problem_sets( + initial_state=[("J/psi(1S)", [-1, 1])], + final_state=["gamma", "pi0", "pi0"], + particle_db=particle_db, + allowed_intermediate_particles=["f(0)(980)", "f(0)(1500)"], + interaction_config=InteractionConfig( + type_settings=create_interaction_settings( + "helicity", particle_db=particle_db, max_angular_momentum=2 + ), + allowed_types=[InteractionType.STRONG], + ), + ) + stripped = strip_spin_projections(collection) + assert isinstance(stripped, QNProblemSetCollection) + n_original = sum(map(len, collection.problem_sets.values())) + n_stripped = sum(map(len, stripped.problem_sets.values())) + assert n_stripped < n_original + + qn_transitions = find_qn_transitions(stripped) + assert len(qn_transitions) > 0 + qn_names = { + qn_type.__name__ + for transition in qn_transitions + for prop_map in [*transition.states.values(), *transition.interactions.values()] + for qn_type in prop_map + } + assert "spin_projection" not in qn_names + assert {"spin_magnitude", "parity", "l_magnitude", "s_magnitude"} <= qn_names + + intermediate_signatures = { + ( + state[EdgeQuantumNumbers.spin_magnitude], + int(state[EdgeQuantumNumbers.parity]), + int(state[EdgeQuantumNumbers.c_parity]), + ) + for transition in qn_transitions + for state in transition.intermediate_states.values() + } + assert intermediate_signatures == {(0, +1, +1)} # both f0 resonances are 0^{++} + collapsed = { + transition.convert(interaction_converter=lambda _: None) + for transition in qn_transitions + } + assert len(collapsed) == 1 + + serialized = json.dumps(asdict(qn_transitions[0])) + assert '"spin_projection"' not in serialized + assert '"spin_magnitude"' in serialized