Construct amplitude models

\(X\to \pi^-\pi^+\pi^-\)

Model definition: x2pipipi-compass-1391643.json.

This page demonstrates deserialization and evaluation of a partial-wave analysis of a diffractively produced \(3\pi\) system. The analysis is performed independently in bins of \(3\pi\) mass and transferred momentum. The full analysis contains about 170 decay chains (about 88 waves, symmetrized for the two \(\pi^+\pi^-\) pairs) per bin. See INSPIRE-HEP 1391643 for details.

Import Python libraries
from copy import deepcopy
from pathlib import Path
import json
import logging

import jax.numpy as jnp
import matplotlib.pyplot as plt
import pandas as pd
import sympy as sp
from ampform.dynamics.form_factor import FormFactor
from ampform.dynamics.phasespace import PhaseSpaceFactorComplex
from ampform_dpd import DefinedExpression
from ampform_dpd.io import cached, perform_cached_lambdify
from ampform_dpd.io.serialization import amplitude as serialization_amplitude
from ampform_dpd.io.serialization.amplitude import formulate
from ampform_dpd.io.serialization.decay import get_final_state, to_decay
from ampform_dpd.io.serialization.dynamics import (
    formulate_dynamics,
    formulate_multichannel_breit_wigner,
    to_mandelstam_symbol,
)
from ampform_dpd.io.serialization.format import (
    get_decay_chains,
    get_function_definition,
)
from matplotlib_inline.backend_inline import set_matplotlib_formats
from sympy.parsing.sympy_parser import parse_expr

THIS_DIR = Path(".").absolute()
logging.getLogger("ampform.sympy").setLevel(logging.ERROR)
set_matplotlib_formats("svg")
with open(THIS_DIR.parent.parent / "models" / "x2pipipi-compass-1391643.json") as f:
    MODEL_DEFINITION = json.load(f)

The JSON file contains four distributions from the same mass bin. Since the Python deserializer constructs one distribution at a time, each distribution is selected and formulated separately.

Compatibility adapters for LS-coupled models
def select_distribution(model_definition: dict, distribution: dict) -> dict:
    selected = deepcopy(model_definition)
    selected["distributions"] = [deepcopy(distribution)]
    for chain_idx, chain in enumerate(get_decay_chains(selected)):
        for vertex in chain["vertices"]:
            if "helicities" in vertex:
                continue
            has_isobar = any(isinstance(item, list) for item in vertex["node"])
            vertex["helicities"] = [str(chain_idx), "0"] if has_isobar else ["0", "0"]
    return selected


def get_ls_child_spins(model_definition, chain_idx, vertex_idx):
    chain = get_decay_chains(model_definition)[chain_idx]
    vertex = chain["vertices"][vertex_idx]
    final_state = get_final_state(model_definition)
    resonance_spin = sp.Rational(chain["propagators"][0]["spin"])
    return tuple(
        final_state[item].spin if isinstance(item, int) else resonance_spin
        for item in vertex["node"]
    )


def formulate_generic_function(propagator, resonance, model):
    definition = get_function_definition(propagator["parametrization"], model)
    expression = definition["expression"]
    expression = expression.replace("^", "**").replace("1im", "I")
    expression = expression.replace("m_12_sq", "sigma")
    mandelstam = to_mandelstam_symbol(propagator["node"])
    return DefinedExpression(
        expression=parse_expr(
            expression,
            local_dict={"i": sp.I, "I": sp.I, "sigma": mandelstam},
        )
    )


def formulate_analytic_multichannel_breit_wigner(propagator, resonance, model):
    dynamics = formulate_multichannel_breit_wigner(propagator, resonance, model)
    s, mass, channels = dynamics.expression.args
    channel_terms = (
        channel.coupling_squared
        * PhaseSpaceFactorComplex(channel.s, channel.m1, channel.m2)
        * FormFactor(
            channel.s,
            channel.m1,
            channel.m2,
            channel.angular_momentum,
            channel.meson_radius,
        )
        ** 2
        for channel in channels
    )
    expression = 1 / (mass**2 - s - sp.I * sum(channel_terms))
    return DefinedExpression(expression=expression, parameters=dynamics.parameters)


serialization_amplitude._get_child_spins = get_ls_child_spins
ADDITIONAL_BUILDERS = {
    "generic_function": formulate_generic_function,
    "MultichannelBreitWigner": formulate_analytic_multichannel_breit_wigner,
}
SELECTED_DEFINITIONS = {
    distribution["name"]: select_distribution(MODEL_DEFINITION, distribution)
    for distribution in MODEL_DEFINITION["distributions"]
}
MODELS = {
    name: formulate(
        definition,
        cleanup_summations=True,
        additional_builders=ADDITIONAL_BUILDERS,
    )
    for name, definition in SELECTED_DEFINITIONS.items()
}
pd.DataFrame({
    "Distribution": list(MODELS),
    "Decay chains": [len(get_decay_chains(model)) for model in SELECTED_DEFINITIONS.values()],
})
Distribution Decay chains
0 compass_3pi_JP=1+_M=0_1540_1560 16
1 compass_3pi_JP=1-_M=1_1540_1560 2
2 compass_3pi_JP=2+_M=1_1540_1560 6
3 compass_3pi_JP=4+_M=1_1540_1560 4

Validation

The table compares all serialized checkpoints with the Python implementation. The marks โ€œ๐ŸŸขโ€, โ€œ๐ŸŸกโ€, and โ€œ๐Ÿ”ดโ€ indicate an absolute difference of \(<10^{-10}\), \(<10^{-2}\), or \(\ge10^{-2}\), respectively. The numerical results are shown even when a checkpoint does not agree with the reference value.

Validation helpers
checksum_points = {
    point["name"]: {parameter["name"]: parameter["value"] for parameter in point["parameters"]}
    for point in MODEL_DEFINITION["parameter_points"]
}


def to_number(value: float | str) -> float | complex:
    if isinstance(value, str):
        value = complex(value.replace(" ", "").replace("i", "j"))
    number = complex(value)
    return number.real if number.imag == 0 else number


def round_number(value: float | complex, digits: int = 6) -> float | complex:
    number = complex(value)
    if number.imag == 0:
        return round(number.real, digits)
    return complex(round(number.real, digits), round(number.imag, digits))


def label_diff(difference: complex) -> str:
    absolute_difference = abs(difference)
    if absolute_difference < 1e-10:
        return "๐ŸŸข"
    if absolute_difference < 1e-2:
        return "๐ŸŸก"
    return "๐Ÿ”ด"


def create_intensity_function(model, subsystem: int):
    invariant = next(s for s in model.invariants if str(s) == f"sigma{subsystem}")
    intensity_expr = cached.xreplace(cached.unfold(model), model.variables)
    intensity_expr = cached.xreplace(intensity_expr, model.parameter_defaults)
    invariant_expr = model.invariants[invariant].xreplace(model.masses).doit()
    intensity_expr = cached.doit(intensity_expr.xreplace({invariant: invariant_expr}))
    return cached.lambdify(intensity_expr, backend="jax")


def evaluate_intensity(model, point: dict[str, float], pair: tuple[int, int]) -> float:
    i, j = pair
    k, *_ = {1, 2, 3} - set(pair)
    invariants = {str(s): s for s in model.invariants}
    sigma_j = invariants[f"sigma{j}"]
    sigma_k = invariants[f"sigma{k}"]
    angle_expr = next(
        expression
        for angle, expression in model.variables.items()
        if str(angle) == f"theta_{i}{j}"
    )
    sigma_k_value = point[f"m_{i}{j}"] ** 2
    equation = sp.Eq(
        sp.cos(angle_expr).xreplace(model.parameter_defaults),
        point[f"cos_theta_{i}{j}"],
    )
    sigma_j_value = sp.solve(equation.subs(sigma_k, sigma_k_value), sigma_j)[0]
    function = create_intensity_function(model, subsystem=i)
    value = function({str(sigma_j): float(sigma_j_value), str(sigma_k): sigma_k_value})
    return float(jnp.real(value))
Compute every serialized checkpoint
chains_by_propagator = {
    propagator["parametrization"]: chain
    for definition in SELECTED_DEFINITIONS.values()
    for chain in get_decay_chains(definition)
    for propagator in chain["propagators"]
}
validation_results = []
for checksum in MODEL_DEFINITION["misc"]["amplitude_model_checksums"]:
    distribution_name = checksum["distribution"]
    point = checksum_points[checksum["point"]]
    reference_value = to_number(checksum["value"])
    chain = chains_by_propagator.get(distribution_name)
    if chain is None:
        computed_value = evaluate_intensity(
            MODELS[distribution_name],
            point,
            pair=(2, 3),
        )
    else:
        dynamics = formulate_dynamics(
            chain,
            MODEL_DEFINITION,
            additional_definitions=ADDITIONAL_BUILDERS,
        )
        expression = dynamics.expression.doit()
        variables = expression.free_symbols - dynamics.parameters.keys()
        if variables:
            function = perform_cached_lambdify(
                expression,
                parameters=dynamics.parameters,
            )
            variable = next(iter(variables))
            computed_value = function({str(variable): next(iter(point.values()))})
        else:
            computed_value = complex(expression.xreplace(dynamics.parameters))
    difference = abs(reference_value - computed_value)
    validation_results.append({
        "Distribution": distribution_name,
        "Point": checksum["point"],
        "Reference": round_number(reference_value),
        "Computed": round_number(computed_value),
        "Difference": difference,
        "Status": label_diff(reference_value - computed_value),
    })

pd.DataFrame(validation_results, dtype=object)
Distribution Point Reference Computed Difference Status
0 compass_3pi_JP=1+_M=0_1540_1560 validation_point1 0.557753 0.557177 0.000577 ๐ŸŸก
1 compass_3pi_JP=1+_M=0_1540_1560 validation_point2 1.436252 1.435651 0.000601 ๐ŸŸก
2 compass_3pi_JP=1+_M=0_1540_1560 validation_point3 0.108395 0.108572 0.000177 ๐ŸŸก
3 compass_3pi_JP=1+_M=0_1540_1560 validation_point4 0.791847 0.792119 0.000272 ๐ŸŸก
4 compass_3pi_JP=1-_M=1_1540_1560 validation_point1 0.045205 0.045205 0.0 ๐ŸŸข
5 compass_3pi_JP=1-_M=1_1540_1560 validation_point2 0.010651 0.010651 0.0 ๐ŸŸข
6 compass_3pi_JP=1-_M=1_1540_1560 validation_point3 0.01857 0.01857 0.0 ๐ŸŸข
7 compass_3pi_JP=1-_M=1_1540_1560 validation_point4 0.090084 0.090084 0.0 ๐ŸŸข
8 compass_3pi_JP=2+_M=1_1540_1560 validation_point1 0.017082 0.017082 0.0 ๐ŸŸข
9 compass_3pi_JP=2+_M=1_1540_1560 validation_point2 0.057953 0.057953 0.0 ๐ŸŸข
10 compass_3pi_JP=2+_M=1_1540_1560 validation_point3 0.035384 0.035384 0.0 ๐ŸŸข
11 compass_3pi_JP=2+_M=1_1540_1560 validation_point4 0.112218 0.112218 0.0 ๐ŸŸข
12 compass_3pi_JP=4+_M=1_1540_1560 validation_point1 0.020556 0.020556 0.0 ๐ŸŸข
13 compass_3pi_JP=4+_M=1_1540_1560 validation_point2 0.011818 0.011818 0.0 ๐ŸŸข
14 compass_3pi_JP=4+_M=1_1540_1560 validation_point3 0.01755 0.01755 0.0 ๐ŸŸข
15 compass_3pi_JP=4+_M=1_1540_1560 validation_point4 0.105833 0.105833 0.0 ๐ŸŸข
16 R(1274) validation_point_m12sq (1.682682+0.620869j) (1.682682+0.620869j) 8.671119018262734e-16 ๐ŸŸข
17 R(1690) validation_point_m12sq (0.564308+0.103183j) (0.564308+0.103183j) 1.2412670766236366e-16 ๐ŸŸข
18 R(600) validation_point_m12sq (0.101724+1.027344j) (0.101724+1.027344j) 2.423651445728339e-16 ๐ŸŸข
19 R(768) validation_point_m12sq (-1.732183+0.632399j) (-1.732183+0.632399j) 3.155846784686366e-15 ๐ŸŸข
20 R(965) validation_point_m12sq (-0.921258+2.14704j) (-0.867171+2.094289j) 0.07555143659694909 ๐Ÿ”ด

Visualization

Dalitz plots

The four distributions in the selected \(3\pi\) mass bin are shown over the same pair of Mandelstam variables.

Compute the four Dalitz plots
x_subsystem, y_subsystem = 3, 1
eliminated_subsystem, *_ = {1, 2, 3} - {x_subsystem, y_subsystem}
resolution = 200
intensity_grids = {}
for name, model in MODELS.items():
    masses = sorted(model.masses, key=str)

    def invariant_limits(subsystem: int) -> tuple[float, float]:
        a, b = sorted({1, 2, 3} - {subsystem})
        minimum = (masses[a] + masses[b]) ** 2
        maximum = (masses[0] - masses[subsystem]) ** 2
        return (
            float(minimum.xreplace(model.masses)),
            float(maximum.xreplace(model.masses)),
        )

    x_min, x_max = invariant_limits(x_subsystem)
    y_min, y_max = invariant_limits(y_subsystem)
    x_margin = 0.05 * (x_max - x_min)
    y_margin = 0.05 * (y_max - y_min)
    X, Y = jnp.meshgrid(
        jnp.linspace(x_min - x_margin, x_max + x_margin, resolution),
        jnp.linspace(y_min - y_margin, y_max + y_margin, resolution),
    )
    function = create_intensity_function(model, subsystem=eliminated_subsystem)
    values = jnp.real(
        function({f"sigma{x_subsystem}": X, f"sigma{y_subsystem}": Y})
    )
    intensity_grids[name] = (X, Y, values / jnp.nansum(values))

# cspell:ignore bbox edgecolor facecolor fontset fontsize labelcolor mathtext ncol
# cspell:ignore hspace savefig sharey STIX wspace xtick ytick
plt.rcParams.update({
    "axes.edgecolor": "#808080",
    "axes.facecolor": "none",
    "axes.labelcolor": "#808080",
    "figure.facecolor": "none",
    "font.family": "serif",
    "font.serif": ["STIXGeneral", "DejaVu Serif"],
    "font.size": 18,
    "mathtext.fontset": "stix",
    "savefig.transparent": True,
    "text.color": "#808080",
    "xtick.color": "#808080",
    "ytick.color": "#808080",
})
Render the four Dalitz plots
subplot_titles = {
    "compass_3pi_JP=1+_M=0_1540_1560": R"$J^P=1^+,\ M=0$",
    "compass_3pi_JP=1-_M=1_1540_1560": R"$J^P=1^-,\ M=1$",
    "compass_3pi_JP=2+_M=1_1540_1560": R"$J^P=2^+,\ M=1$",
    "compass_3pi_JP=4+_M=1_1540_1560": R"$J^P=4^+,\ M=1$",
}
fig, axes = plt.subplots(2, 2, figsize=(14, 11))
for ax, (name, (X, Y, values)) in zip(axes.flat, intensity_grids.items(), strict=True):
    ax.pcolormesh(X, Y, values, rasterized=True)
    ax.set_aspect("equal")
    ax.set_title(subplot_titles[name])
    ax.set_xlabel(R"$m^2(\pi^-\pi^+)$ [GeV$^2$]")
    ax.set_ylabel(R"$m^2(\pi^+\pi^-)$ [GeV$^2$]")
fig.subplots_adjust(hspace=0.38, wspace=0.08)
plt.show()