Skip to content

kups.potential.common.pair

Composable pair energies with independent features, cutoffs, and masks.

HasPairTerms

Bases: Protocol

Flat terms exposed by a single pair energy or a sum for composition.

Source code in src/kups/potential/common/pair.py
@runtime_checkable
class HasPairTerms[Params, Part, Feat](Protocol):
    """Flat terms exposed by a single pair energy or a sum for composition."""

    @property
    def terms(self) -> tuple[PairTerm[Params, Part, Feat], ...]: ...

PairBatch

Bases: NamedTuple

Geometry and independent masks for a batch of candidate pairs.

valid checks active endpoints and matching systems. Inclusion and exclusion policies are kept separate so several consumers can share the same candidates. Vectors point from query to key. Invalid geometry is sanitized before a consumer evaluates singular pair kernels.

Source code in src/kups/potential/common/pair.py
class PairBatch(NamedTuple):
    """Geometry and independent masks for a batch of candidate pairs.

    ``valid`` checks active endpoints and matching systems. Inclusion and
    exclusion policies are kept separate so several consumers can share the
    same candidates. Vectors point from query to key. Invalid geometry is
    sanitized before a consumer evaluates singular pair kernels.
    """

    rij: Array
    r2: Array
    system: Index[SystemId]
    valid: Array
    inclusion: Array
    exclusion: Array

    @classmethod
    def from_candidates(
        cls,
        batch: CandidateBatch[Literal[2]],
        ctx: PipelineContext,
        *,
        query_lanes: int | None = None,
    ) -> Self:
        """Prepare pair geometry and masks from any selector's candidate batch.

        ``query_lanes`` declares a regular block of candidates for each query, in
        query-table order. This shares one cell matrix across a whole lane block
        and broadcasts query coordinates so their gradients reduce along lanes.
        """
        keys = ctx.keys[batch.key_idx]
        queries = (
            ctx.edge_query_table[batch.query_idx]
            if query_lanes is None
            else jax.tree.map(
                lambda x: jnp.repeat(x, query_lanes, axis=0),
                ctx.edge_query_table.data,
            )
        )
        key_system, query_system = Index.match(keys.system, queries.system)
        valid = InBoundsMask()(batch, ctx) & (key_system == query_system)
        delta = keys.positions - queries.positions - batch.edges.shifts[:, 0]
        # Avoid inf - inf propagating NaNs into cell/position derivatives.
        delta = jnp.where(valid[:, None], delta, 0.0)
        frames = ctx.systems.map_data(lambda s: s.cell.frame.materialize())
        if query_lanes is None:
            rij = frames[keys.system].to_real(delta)
        else:
            query_table = ctx.edge_query_table
            delta = delta.reshape(query_table.size, query_lanes, 3)
            # Keep the frame's diagonal/triangular structure and share it across
            # lanes. A dense batched matmul can become a separate GPU BLAS call,
            # materializing pair vectors instead of fusing with the pair kernel.
            query_frames = jax.tree.map(
                lambda x: x[:, None], frames[query_table.data.system]
            )
            rij = query_frames.to_real(delta).reshape(-1, 3)
        r2 = jnp.sum(rij**2, axis=-1)
        return cls(
            rij,
            r2,
            keys.system,
            valid,
            InclusionMatchMask()(batch, ctx),
            ExclusionMask()(batch, ctx),
        )

from_candidates(batch, ctx, *, query_lanes=None) classmethod

Prepare pair geometry and masks from any selector's candidate batch.

query_lanes declares a regular block of candidates for each query, in query-table order. This shares one cell matrix across a whole lane block and broadcasts query coordinates so their gradients reduce along lanes.

Source code in src/kups/potential/common/pair.py
@classmethod
def from_candidates(
    cls,
    batch: CandidateBatch[Literal[2]],
    ctx: PipelineContext,
    *,
    query_lanes: int | None = None,
) -> Self:
    """Prepare pair geometry and masks from any selector's candidate batch.

    ``query_lanes`` declares a regular block of candidates for each query, in
    query-table order. This shares one cell matrix across a whole lane block
    and broadcasts query coordinates so their gradients reduce along lanes.
    """
    keys = ctx.keys[batch.key_idx]
    queries = (
        ctx.edge_query_table[batch.query_idx]
        if query_lanes is None
        else jax.tree.map(
            lambda x: jnp.repeat(x, query_lanes, axis=0),
            ctx.edge_query_table.data,
        )
    )
    key_system, query_system = Index.match(keys.system, queries.system)
    valid = InBoundsMask()(batch, ctx) & (key_system == query_system)
    delta = keys.positions - queries.positions - batch.edges.shifts[:, 0]
    # Avoid inf - inf propagating NaNs into cell/position derivatives.
    delta = jnp.where(valid[:, None], delta, 0.0)
    frames = ctx.systems.map_data(lambda s: s.cell.frame.materialize())
    if query_lanes is None:
        rij = frames[keys.system].to_real(delta)
    else:
        query_table = ctx.edge_query_table
        delta = delta.reshape(query_table.size, query_lanes, 3)
        # Keep the frame's diagonal/triangular structure and share it across
        # lanes. A dense batched matmul can become a separate GPU BLAS call,
        # materializing pair vectors instead of fusing with the pair kernel.
        query_frames = jax.tree.map(
            lambda x: x[:, None], frames[query_table.data.system]
        )
        rij = query_frames.to_real(delta).reshape(-1, 3)
    r2 = jnp.sum(rij**2, axis=-1)
    return cls(
        rij,
        r2,
        keys.system,
        valid,
        InclusionMatchMask()(batch, ctx),
        ExclusionMask()(batch, ctx),
    )

PairEnergy

A pair kernel with feature, cutoff, and neighbor-mask policies.

with_parameters(view) binds a term to a field of shared parameters. For example, lj.with_parameters(lambda p: p.lj) + coulomb.with_parameters(lambda p: p.ewald) uses one neighbor traversal at the maximum cutoff, retaining each term's individual cutoff. Terms use a shared particle type; with_particles(view) adapts each feature selector to that type before adding them.

Set inclusion=False to interact across inclusion groups within a system, and exclusion=False to include same-group pairs. Exclusions apply to the closest image; other periodic copies remain eligible. Zero-shift self pairs and inactive particles are always excluded.

Source code in src/kups/potential/common/pair.py
@dataclass
class PairEnergy[Params, Part, Feat]:
    """A pair kernel with feature, cutoff, and neighbor-mask policies.

    ``with_parameters(view)`` binds a term to a field of shared parameters.
    For example, ``lj.with_parameters(lambda p: p.lj) +
    coulomb.with_parameters(lambda p: p.ewald)`` uses one neighbor traversal
    at the maximum cutoff, retaining each term's individual cutoff.
    Terms use a shared particle type; ``with_particles(view)`` adapts each
    feature selector to that type before adding them.

    Set ``inclusion=False`` to interact across inclusion groups within a
    system, and ``exclusion=False`` to include same-group pairs. Exclusions
    apply to the closest image; other periodic copies remain eligible.
    Zero-shift self pairs and inactive particles are always excluded.
    """

    kernel: PairKernel[Params, Feat] = field(static=True)
    features: View[Part, Feat] = field(static=True)
    cutoffs: View[Params, Table[SystemId, Array]] = field(static=True)
    inclusion: bool = field(static=True, default=True)
    exclusion: bool = field(static=True, default=True)

    @property
    def terms(self) -> tuple[PairTerm[Params, Part, Feat]]:
        """Expose this energy as one term for flat composition."""
        return (self,)

    def __add__[Other](
        self,
        other: HasPairTerms[Params, Part, Other],
    ) -> PairEnergySum[Params, Part, Feat | Other]:
        return _add(self, other)

    def evaluate(
        self,
        parameters: Params,
        left: Feat,
        right: Feat,
        pairs: PairBatch,
    ) -> Array:
        cutoffs = self.cutoffs(parameters)
        cutoff = cutoffs.data[0] if cutoffs.size == 1 else cutoffs[pairs.system]
        keep = pairs.valid & (pairs.r2 < cutoff**2)
        if self.inclusion:
            keep &= pairs.inclusion
        if self.exclusion:
            keep &= pairs.exclusion
        rij = jnp.where(keep[..., None], pairs.rij, 1.0)
        r2 = jnp.where(keep, pairs.r2, 3.0)
        return jnp.where(
            keep,
            self.kernel(
                parameters,
                left,
                right,
                rij,
                r2,
                pairs.system,
            ),
            0.0,
        )

    def with_particles[Outer](
        self, view: View[Outer, Part]
    ) -> PairEnergy[Params, Outer, Feat]:
        """Select this term's particle interface from a shared particle type."""
        return PairEnergy(
            self.kernel,
            pipe(view, self.features),
            self.cutoffs,
            self.inclusion,
            self.exclusion,
        )

    def with_parameters[Outer](
        self,
        view: View[Outer, Params],
    ) -> PairEnergy[Outer, Part, Feat]:
        """Read this term's parameters from a shared parameter bundle."""

        return PairEnergy(
            lambda parameters, left, right, rij, r2, system: self.kernel(
                view(parameters), left, right, rij, r2, system
            ),
            self.features,
            pipe(view, self.cutoffs),
            self.inclusion,
            self.exclusion,
        )

terms property

Expose this energy as one term for flat composition.

with_parameters(view)

Read this term's parameters from a shared parameter bundle.

Source code in src/kups/potential/common/pair.py
def with_parameters[Outer](
    self,
    view: View[Outer, Params],
) -> PairEnergy[Outer, Part, Feat]:
    """Read this term's parameters from a shared parameter bundle."""

    return PairEnergy(
        lambda parameters, left, right, rij, r2, system: self.kernel(
            view(parameters), left, right, rij, r2, system
        ),
        self.features,
        pipe(view, self.cutoffs),
        self.inclusion,
        self.exclusion,
    )

with_particles(view)

Select this term's particle interface from a shared particle type.

Source code in src/kups/potential/common/pair.py
def with_particles[Outer](
    self, view: View[Outer, Part]
) -> PairEnergy[Params, Outer, Feat]:
    """Select this term's particle interface from a shared particle type."""
    return PairEnergy(
        self.kernel,
        pipe(view, self.features),
        self.cutoffs,
        self.inclusion,
        self.exclusion,
    )

PairEnergySum

A flat sum of pair terms evaluated on the same candidate batch.

Use + to combine terms with distinct feature types. Each entry in the feature tuple comes from, and is passed back to, the term at the same position. Feat retains the union of the terms' concrete feature types. from_term starts a sum from any implementation of PairTerm.

Source code in src/kups/potential/common/pair.py
@dataclass
class PairEnergySum[Params, Part, Feat]:
    """A flat sum of pair terms evaluated on the same candidate batch.

    Use ``+`` to combine terms with distinct feature types. Each entry in the
    feature tuple comes from, and is passed back to, the term at the same
    position. ``Feat`` retains the union of the terms' concrete feature types.
    ``from_term`` starts a sum from any implementation of ``PairTerm``.
    """

    terms: tuple[PairTerm[Params, Part, Feat], ...] = field(static=True)

    def __post_init__(self) -> None:
        if not self.terms:
            raise ValueError("At least one pair energy is required")
        for i, term in enumerate(self.terms):
            if not isinstance(term, PairTerm):
                raise TypeError(
                    f"PairEnergySum term {i} must implement PairTerm; "
                    f"got {type(term).__name__}"
                )

    @classmethod
    def from_term(cls, term: PairTerm[Params, Part, Feat]) -> Self:
        """Start a sum with the term's concrete feature type."""
        return cls((term,))

    def __add__[Other](
        self,
        other: HasPairTerms[Params, Part, Other],
    ) -> PairEnergySum[Params, Part, Feat | Other]:
        return _add(self, other)

    def features(self, particles: Part, /) -> tuple[Feat, ...]:
        return tuple(term.features(particles) for term in self.terms)

    @property
    def inclusion(self) -> bool:
        return all(term.inclusion for term in self.terms)

    @property
    def exclusion(self) -> bool:
        return all(term.exclusion for term in self.terms)

    def cutoffs(self, parameters: Params, /) -> Table[SystemId, Array]:
        tables = [term.cutoffs(parameters) for term in self.terms]
        # A singleton cutoff broadcasts to every system, as elsewhere in kUPS.
        target = max(tables, key=lambda t: t.size)
        return target.set_data(
            jnp.stack([Table.broadcast_to(t, target).data for t in tables]).max(axis=0)
        )

    def evaluate(
        self,
        parameters: Params,
        left: tuple[Feat, ...],
        right: tuple[Feat, ...],
        pairs: PairBatch,
    ) -> Array:
        values = [
            term.evaluate(parameters, a, b, pairs)
            for term, a, b in zip(self.terms, left, right, strict=True)
        ]
        return sum(values[1:], values[0])

from_term(term) classmethod

Start a sum with the term's concrete feature type.

Source code in src/kups/potential/common/pair.py
@classmethod
def from_term(cls, term: PairTerm[Params, Part, Feat]) -> Self:
    """Start a sum with the term's concrete feature type."""
    return cls((term,))

PairKernel

Bases: Protocol

Numerical energy formula evaluated on already selected pairs.

A kernel takes parameters, the two endpoints' features and their geometry, and returns one energy per pair. For example, the Lennard-Jones kernel uses species labels to look up mixing parameters and evaluates the 12-6 formula. It leaves neighbor selection, cutoffs, masks and summation to its callers. In particular, it does not halve energies to account for directed edges.

Use JAX-compatible array operations and support broadcasting over the pair axes: the same kernel handles flat graph edges and rectangular query/key blocks. PairEnergy replaces masked pairs' geometry with finite, nonzero values before calling the kernel, then sets their returned energies to zero.

Class Type Parameters:

Name Bound or Constraints Description Default
Params

The term's parameter bundle, such as mixing tables or screening constants.

required
Feat

Per-particle feature pytree selected by PairTerm.features, such as species indices or charges, gathered for each endpoint.

required
Source code in src/kups/potential/common/pair.py
class PairKernel[Params, Feat](Protocol):
    """Numerical energy formula evaluated on already selected pairs.

    A kernel takes parameters, the two endpoints' features and their geometry,
    and returns one energy per pair. For example, the Lennard-Jones kernel uses
    species labels to look up mixing parameters and evaluates the 12-6 formula.
    It leaves neighbor selection, cutoffs, masks and summation to its callers.
    In particular, it does not halve energies to account for directed edges.

    Use JAX-compatible array operations and support broadcasting over the pair
    axes: the same kernel handles flat graph edges and rectangular query/key
    blocks. ``PairEnergy`` replaces masked pairs' geometry with finite, nonzero
    values before calling the kernel, then sets their returned energies to zero.

    Type Parameters:
        Params: The term's parameter bundle, such as mixing tables or screening
            constants.
        Feat: Per-particle feature pytree selected by ``PairTerm.features``, such
            as species indices or charges, gathered for each endpoint.
    """

    def __call__(
        self,
        parameters: Params,
        features_i: Feat,
        features_j: Feat,
        rij: Array,
        r2: Array,
        system: Index[SystemId],
        /,
    ) -> Array:
        """Evaluate the interaction without reducing its pair axes.

        Args:
            parameters: Parameters for this interaction.
            features_i: Left-endpoint features, broadcastable over the pair axes.
            features_j: Right-endpoint features, with the same pytree structure.
            rij: Displacement from left to right, including the periodic shift,
                with shape ``(*pair_shape, 3)``.
            r2: Squared distances, with shape ``pair_shape``.
            system: System ids for selecting per-system parameters.

        Returns:
            Energies with shape ``pair_shape``, in the potential's energy units.
        """
        ...

__call__(parameters, features_i, features_j, rij, r2, system)

Evaluate the interaction without reducing its pair axes.

Parameters:

Name Type Description Default
parameters Params

Parameters for this interaction.

required
features_i Feat

Left-endpoint features, broadcastable over the pair axes.

required
features_j Feat

Right-endpoint features, with the same pytree structure.

required
rij Array

Displacement from left to right, including the periodic shift, with shape (*pair_shape, 3).

required
r2 Array

Squared distances, with shape pair_shape.

required
system Index[SystemId]

System ids for selecting per-system parameters.

required

Returns:

Type Description
Array

Energies with shape pair_shape, in the potential's energy units.

Source code in src/kups/potential/common/pair.py
def __call__(
    self,
    parameters: Params,
    features_i: Feat,
    features_j: Feat,
    rij: Array,
    r2: Array,
    system: Index[SystemId],
    /,
) -> Array:
    """Evaluate the interaction without reducing its pair axes.

    Args:
        parameters: Parameters for this interaction.
        features_i: Left-endpoint features, broadcastable over the pair axes.
        features_j: Right-endpoint features, with the same pytree structure.
        rij: Displacement from left to right, including the periodic shift,
            with shape ``(*pair_shape, 3)``.
        r2: Squared distances, with shape ``pair_shape``.
        system: System ids for selecting per-system parameters.

    Returns:
        Energies with shape ``pair_shape``, in the potential's energy units.
    """
    ...

PairTerm

Bases: Protocol

Complete pair interaction interface consumed by neighbor evaluators.

A term selects particle features, declares the search cutoff, and evaluates candidate pairs with its own cutoff and mask policy. PairEnergy implements this interface for one PairKernel. PairEnergySum combines terms: its search cutoff is their maximum, while evaluation retains each term's cutoff. Its inclusion/exclusion flags describe masks required by every contribution, allowing neighbor selection to apply those shared masks before compaction.

The evaluator constructs neighbors and periodic geometry, gathers endpoint features, and reduces the returned pair energies into system totals. GraphPairEnergy derives a graph evaluator from a PairEnergy; FusedNeighborEnergy accepts any PairTerm for shared full/local evaluation.

Class Type Parameters:

Name Bound or Constraints Description Default
Params

Parameters supplied to the term's cutoff and energy functions.

required
Part

Particle data presented to the feature selector.

required
Feat

Selected per-particle feature pytree. Feature arrays preserve the leading particle axis; evaluate receives gathered pair endpoints.

required
Source code in src/kups/potential/common/pair.py
@runtime_checkable
class PairTerm[Params, Part, Feat](Protocol):
    """Complete pair interaction interface consumed by neighbor evaluators.

    A term selects particle features, declares the search cutoff, and evaluates
    candidate pairs with its own cutoff and mask policy. ``PairEnergy`` implements
    this interface for one ``PairKernel``. ``PairEnergySum`` combines terms: its
    search cutoff is their maximum, while evaluation retains each term's cutoff.
    Its inclusion/exclusion flags describe masks required by every contribution,
    allowing neighbor selection to apply those shared masks before compaction.

    The evaluator constructs neighbors and periodic geometry, gathers endpoint
    features, and reduces the returned pair energies into system totals.
    ``GraphPairEnergy`` derives a graph evaluator from a ``PairEnergy``;
    ``FusedNeighborEnergy`` accepts any ``PairTerm`` for shared full/local evaluation.

    Type Parameters:
        Params: Parameters supplied to the term's cutoff and energy functions.
        Part: Particle data presented to the feature selector.
        Feat: Selected per-particle feature pytree. Feature arrays preserve the
            leading particle axis; ``evaluate`` receives gathered pair endpoints.
    """

    @property
    def features(self) -> View[Part, Feat]:
        """Select the payload needed by the kernel, e.g. labels or charges.

        Evaluators may cache these per-particle features in their cell table.
        """
        ...

    @property
    def cutoffs(self) -> View[Params, Table[SystemId, Array]]:
        """Select per-system search radii; a singleton table applies to all systems."""
        ...

    @property
    def inclusion(self) -> bool:
        """Whether every contribution requires matching inclusion groups."""
        ...

    @property
    def exclusion(self) -> bool:
        """Whether every contribution excludes same-group closest images."""
        ...

    def evaluate(
        self,
        parameters: Params,
        left: Feat,
        right: Feat,
        pairs: PairBatch,
    ) -> Array:
        """Return pair energies after applying this term's cutoff and masks.

        ``left`` and ``right`` contain endpoint features broadcastable over
        ``pairs.r2.shape``. The result has that shape, with zero contributions
        for invalid or filtered pairs. Geometry and candidate masks come from
        ``PairBatch``; the term decides which inclusion/exclusion masks apply.
        Summation and directed-edge counting remain the evaluator's responsibility.
        """
        ...

cutoffs property

Select per-system search radii; a singleton table applies to all systems.

exclusion property

Whether every contribution excludes same-group closest images.

features property

Select the payload needed by the kernel, e.g. labels or charges.

Evaluators may cache these per-particle features in their cell table.

inclusion property

Whether every contribution requires matching inclusion groups.

evaluate(parameters, left, right, pairs)

Return pair energies after applying this term's cutoff and masks.

left and right contain endpoint features broadcastable over pairs.r2.shape. The result has that shape, with zero contributions for invalid or filtered pairs. Geometry and candidate masks come from PairBatch; the term decides which inclusion/exclusion masks apply. Summation and directed-edge counting remain the evaluator's responsibility.

Source code in src/kups/potential/common/pair.py
def evaluate(
    self,
    parameters: Params,
    left: Feat,
    right: Feat,
    pairs: PairBatch,
) -> Array:
    """Return pair energies after applying this term's cutoff and masks.

    ``left`` and ``right`` contain endpoint features broadcastable over
    ``pairs.r2.shape``. The result has that shape, with zero contributions
    for invalid or filtered pairs. Geometry and candidate masks come from
    ``PairBatch``; the term decides which inclusion/exclusion masks apply.
    Summation and directed-edge counting remain the evaluator's responsibility.
    """
    ...