Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 9 additions & 2 deletions src/qrules/io/_dict.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
62 changes: 62 additions & 0 deletions src/qrules/solving.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
115 changes: 114 additions & 1 deletion src/qrules/workflow.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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,
Expand All @@ -66,6 +68,8 @@
remove_duplicate_solutions,
)
from qrules.topology import (
FrozenDict,
FrozenTransition,
MutableTransition,
create_isobar_topologies,
create_n_body_topology,
Expand Down Expand Up @@ -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],
Expand Down
57 changes: 57 additions & 0 deletions tests/unit/test_workflow.py
Original file line number Diff line number Diff line change
@@ -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 (
Expand All @@ -13,7 +16,9 @@
QNProblemSetCollection,
create_qn_problem_sets,
filter_intermediate_particles,
find_qn_transitions,
find_solutions,
strip_spin_projections,
)


Expand Down Expand Up @@ -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
Loading