Skip to content

torchmodal.nn

torchmodal.nn.operators

torchmodal.nn.operators ~~~~~~~~~~~~~~~~~~~~~~~

Differentiable aggregation operators as nn.Module wrappers.

These modules wrap the functional API in :mod:torchmodal.functional, adding learnable or configurable temperature parameters.

.. note:: Named SmoothMin / SmoothMax (not Softmin / Softmax) to avoid confusion with the standard probability-normalization torch.softmax. Legacy aliases Softmin and Softmax are provided for backward compatibility.

SmoothMin

Bases: Module

Differentiable smooth minimum module (log-sum-exp lower bound).

Sound lower bound on :func:torch.min: smooth_min(x) <= min(x) for x_i \in [0, 1].

Parameters:

Name Type Description Default
tau float

Initial temperature. Default 0.1.

0.1
learnable bool

If True, tau is a learnable parameter.

False
dim int

Dimension to aggregate over. Default -1.

-1
Source code in torchmodal/nn/operators.py
class SmoothMin(nn.Module):
    r"""Differentiable smooth minimum module (log-sum-exp lower bound).

    Sound lower bound on :func:`torch.min`:
    ``smooth_min(x) <= min(x)`` for ``x_i \in [0, 1]``.

    Args:
        tau: Initial temperature. Default 0.1.
        learnable: If ``True``, ``tau`` is a learnable parameter.
        dim: Dimension to aggregate over. Default -1.
    """

    def __init__(
        self,
        tau: float = 0.1,
        learnable: bool = False,
        dim: int = -1,
    ) -> None:
        super().__init__()
        if learnable:
            self.tau = nn.Parameter(torch.tensor(tau))
        else:
            self.register_buffer("tau", torch.tensor(tau))
        self.dim = dim

    def forward(self, x: Tensor) -> Tensor:
        return F.smooth_min(x, tau=self.tau.item(), dim=self.dim)

    def extra_repr(self) -> str:
        return f"tau={self.tau.item():.4f}, dim={self.dim}"

SmoothMax

Bases: Module

Differentiable smooth maximum module (log-sum-exp upper bound).

Sound upper bound on :func:torch.max: smooth_max(x) >= max(x) for x_i \in [0, 1].

Parameters:

Name Type Description Default
tau float

Initial temperature. Default 0.1.

0.1
learnable bool

If True, tau is a learnable parameter.

False
dim int

Dimension to aggregate over. Default -1.

-1
Source code in torchmodal/nn/operators.py
class SmoothMax(nn.Module):
    r"""Differentiable smooth maximum module (log-sum-exp upper bound).

    Sound upper bound on :func:`torch.max`:
    ``smooth_max(x) >= max(x)`` for ``x_i \in [0, 1]``.

    Args:
        tau: Initial temperature. Default 0.1.
        learnable: If ``True``, ``tau`` is a learnable parameter.
        dim: Dimension to aggregate over. Default -1.
    """

    def __init__(
        self,
        tau: float = 0.1,
        learnable: bool = False,
        dim: int = -1,
    ) -> None:
        super().__init__()
        if learnable:
            self.tau = nn.Parameter(torch.tensor(tau))
        else:
            self.register_buffer("tau", torch.tensor(tau))
        self.dim = dim

    def forward(self, x: Tensor) -> Tensor:
        return F.smooth_max(x, tau=self.tau.item(), dim=self.dim)

    def extra_repr(self) -> str:
        return f"tau={self.tau.item():.4f}, dim={self.dim}"

ConvPool

Bases: Module

Convex pooling module.

Parameters:

Name Type Description Default
tau float

Initial temperature. Default 0.1.

0.1
learnable bool

If True, tau is a learnable parameter.

False
dim int

Dimension to pool over. Default -1.

-1
Source code in torchmodal/nn/operators.py
class ConvPool(nn.Module):
    r"""Convex pooling module.

    Args:
        tau: Initial temperature. Default 0.1.
        learnable: If ``True``, ``tau`` is a learnable parameter.
        dim: Dimension to pool over. Default -1.
    """

    def __init__(
        self,
        tau: float = 0.1,
        learnable: bool = False,
        dim: int = -1,
    ) -> None:
        super().__init__()
        if learnable:
            self.tau = nn.Parameter(torch.tensor(tau))
        else:
            self.register_buffer("tau", torch.tensor(tau))
        self.dim = dim

    def forward(self, x: Tensor, z: Tensor | None = None) -> Tensor:
        if z is None:
            z = x
        return F.conv_pool(x, z, tau=self.tau.item(), dim=self.dim)

    def extra_repr(self) -> str:
        return f"tau={self.tau.item():.4f}, dim={self.dim}"

Softmin

Softmin(*args, **kwargs) -> SmoothMin

Deprecated alias for :class:SmoothMin.

Source code in torchmodal/nn/operators.py
def Softmin(*args, **kwargs) -> SmoothMin:  # noqa: N802
    """Deprecated alias for :class:`SmoothMin`."""
    warnings.warn(
        "torchmodal.nn.Softmin is deprecated, use SmoothMin",
        DeprecationWarning,
        stacklevel=2,
    )
    return SmoothMin(*args, **kwargs)

Softmax

Softmax(*args, **kwargs) -> SmoothMax

Deprecated alias for :class:SmoothMax.

Source code in torchmodal/nn/operators.py
def Softmax(*args, **kwargs) -> SmoothMax:  # noqa: N802
    """Deprecated alias for :class:`SmoothMax`."""
    warnings.warn(
        "torchmodal.nn.Softmax is deprecated, use SmoothMax",
        DeprecationWarning,
        stacklevel=2,
    )
    return SmoothMax(*args, **kwargs)

torchmodal.nn.connectives

torchmodal.nn.connectives ~~~~~~~~~~~~~~~~~~~~~~~~~

Propositional logic connectives as nn.Module wrappers.

These implement Łukasiewicz fuzzy logic operators over real-valued truth bounds in [0, 1], following the LNN framework (Riegel et al., 2020) as extended by MLNN (Sulc, 2026).

Each connective operates on truth bounds [L, U] ⊆ [0, 1] and preserves the bound invariant L <= U.

Negation

Bases: Module

Fuzzy negation: :math:\neg x = 1 - x.

For bounds, swaps and negates: [L', U'] = [1-U, 1-L].

Source code in torchmodal/nn/connectives.py
class Negation(nn.Module):
    r"""Fuzzy negation: :math:`\neg x = 1 - x`.

    For bounds, swaps and negates: ``[L', U'] = [1-U, 1-L]``.
    """

    def forward(self, x: Tensor) -> Tensor:
        if x.dim() >= 1 and x.shape[-1] == 2:
            # Bounds tensor: swap L and U after negation
            neg = F.negation(x)
            return neg.flip(-1)
        return F.negation(x)

Conjunction

Bases: Module

Łukasiewicz conjunction (fuzzy AND).

For bounds: - L_{a∧b} = max(0, L_a + L_b - 1) - U_{a∧b} = min(U_a, U_b)

Source code in torchmodal/nn/connectives.py
class Conjunction(nn.Module):
    r"""Łukasiewicz conjunction (fuzzy AND).

    For bounds:
    - ``L_{a∧b} = max(0, L_a + L_b - 1)``
    - ``U_{a∧b} = min(U_a, U_b)``
    """

    def forward(self, a: Tensor, b: Tensor) -> Tensor:
        if a.dim() >= 1 and a.shape[-1] == 2:
            import torch

            L = F.conjunction(a[..., 0], b[..., 0])
            U = torch.min(a[..., 1], b[..., 1])
            return torch.stack([L, U], dim=-1)
        return F.conjunction(a, b)

Disjunction

Bases: Module

Łukasiewicz disjunction (fuzzy OR).

For bounds: - L_{a∨b} = max(L_a, L_b) - U_{a∨b} = min(1, U_a + U_b)

Source code in torchmodal/nn/connectives.py
class Disjunction(nn.Module):
    r"""Łukasiewicz disjunction (fuzzy OR).

    For bounds:
    - ``L_{a∨b} = max(L_a, L_b)``
    - ``U_{a∨b} = min(1, U_a + U_b)``
    """

    def forward(self, a: Tensor, b: Tensor) -> Tensor:
        if a.dim() >= 1 and a.shape[-1] == 2:
            import torch

            L = torch.max(a[..., 0], b[..., 0])
            U = F.disjunction(a[..., 1], b[..., 1])
            return torch.stack([L, U], dim=-1)
        return F.disjunction(a, b)

Implication

Bases: Module

Łukasiewicz implication: :math:a \to b = \min(1, 1 - a + b).

For bounds: - L_{a→b} = max(0, 1 - U_a + L_b) (strongest constraint) - U_{a→b} = min(1, 1 - L_a + U_b)

Source code in torchmodal/nn/connectives.py
class Implication(nn.Module):
    r"""Łukasiewicz implication: :math:`a \to b = \min(1, 1 - a + b)`.

    For bounds:
    - ``L_{a→b} = max(0, 1 - U_a + L_b)``  (strongest constraint)
    - ``U_{a→b} = min(1, 1 - L_a + U_b)``
    """

    def forward(self, a: Tensor, b: Tensor) -> Tensor:
        if a.dim() >= 1 and a.shape[-1] == 2:
            import torch

            L = torch.clamp(1.0 - a[..., 1] + b[..., 0], min=0.0, max=1.0)
            U = torch.clamp(1.0 - a[..., 0] + b[..., 1], min=0.0, max=1.0)
            return torch.stack([L, U], dim=-1)
        return F.implication(a, b)

torchmodal.nn.modal

torchmodal.nn.modal ~~~~~~~~~~~~~~~~~~~

Core modal operator neurons: Necessity (□) and Possibility (♢).

These are the central building blocks of the MLNN framework, implementing differentiable Kripke semantics (Section 3.2.1 of the paper).

The Necessity neuron acts as a "weakest link" detector — aggregating truth values across accessible worlds via differentiable implication.

The Possibility neuron acts as an "evidence scout" — seeking any accessible world where the proposition holds.

Necessity

Bases: Module

Necessity (Box / □) neuron.

Implements the differentiable universal quantification over accessible worlds:

.. math:: L_{\Box\phi,w} = \operatorname{smooth_min}\tau \bigl{ (1 - \tilde{A}{w,w'}) + L_{\phi,w'} \bigr}_{w' \in W}

.. math:: U_{\Box\phi,w} = \operatorname{conv_pool}\tau \bigl{ (1 - \tilde{A}{w,w'}) + U_{\phi,w'} \bigr}_{w' \in W}

With top_k=k each endpoint aggregates only the k smallest of its own implication terms ((1 - Ã) + L for the lower bound, (1 - Ã) + U for the upper). The selection is made on the aggregated terms, not on à alone, so the true minimum is always kept: the bounds stay sound, the smooth lower bound is within tau * log(k) of the crisp minimum, and the result does not depend on |W|. This is where top-k masking belongs — it used to live on the accessibility modules, which was unsound (see :func:torchmodal.functional.necessity).

Parameters:

Name Type Description Default
tau float

Temperature for soft aggregation. Default 0.1.

0.1
learnable_tau bool

If True, temperature is learnable. Default False.

False
top_k Optional[int]

If set, aggregate only the top_k smallest implication terms per world and endpoint. Default None (full row).

None

For temperature annealing during training, update the temperature via :meth:set_tau rather than assigning to .tau (buffers/parameters cannot be assigned a plain float).

Example::

>>> box = torchmodal.nn.Necessity(tau=0.1)
>>> # prop_bounds: (|W|, 2) truth bounds for proposition ϕ
>>> # A: (|W|, |W|) accessibility matrix
>>> box_phi = box(prop_bounds, A)
>>> box.set_tau(0.05)  # annealing
>>> box_k = torchmodal.nn.Necessity(tau=0.1, top_k=8)  # k-neighbourhoods
Source code in torchmodal/nn/modal.py
class Necessity(nn.Module):
    r"""Necessity (Box / □) neuron.

    Implements the differentiable universal quantification over accessible
    worlds:

    .. math::
        L_{\Box\phi,w} = \operatorname{smooth\_min}_\tau \bigl\{
            (1 - \tilde{A}_{w,w'}) + L_{\phi,w'} \bigr\}_{w' \in W}

    .. math::
        U_{\Box\phi,w} = \operatorname{conv\_pool}_\tau \bigl\{
            (1 - \tilde{A}_{w,w'}) + U_{\phi,w'} \bigr\}_{w' \in W}

    With ``top_k=k`` each endpoint aggregates only the *k smallest* of its
    own implication terms (``(1 - Ã) + L`` for the lower bound,
    ``(1 - Ã) + U`` for the upper). The selection is made on the aggregated
    terms, not on ``Ã`` alone, so the true minimum is always kept: the
    bounds stay sound, the smooth lower bound is within ``tau * log(k)``
    of the crisp minimum, and the result does not depend on ``|W|``. This
    is where top-k masking belongs — it used to live on the accessibility
    modules, which was unsound (see :func:`torchmodal.functional.necessity`).

    Args:
        tau: Temperature for soft aggregation. Default 0.1.
        learnable_tau: If ``True``, temperature is learnable. Default False.
        top_k: If set, aggregate only the ``top_k`` smallest implication
            terms per world and endpoint. Default ``None`` (full row).

    For temperature annealing during training, update the temperature via
    :meth:`set_tau` rather than assigning to ``.tau`` (buffers/parameters
    cannot be assigned a plain float).

    Example::

        >>> box = torchmodal.nn.Necessity(tau=0.1)
        >>> # prop_bounds: (|W|, 2) truth bounds for proposition ϕ
        >>> # A: (|W|, |W|) accessibility matrix
        >>> box_phi = box(prop_bounds, A)
        >>> box.set_tau(0.05)  # annealing
        >>> box_k = torchmodal.nn.Necessity(tau=0.1, top_k=8)  # k-neighbourhoods
    """

    def __init__(
        self,
        tau: float = 0.1,
        learnable_tau: bool = False,
        top_k: Optional[int] = None,
    ) -> None:
        super().__init__()
        if learnable_tau:
            self.tau = nn.Parameter(torch.tensor(tau))
        else:
            self.register_buffer("tau", torch.tensor(tau))
        if top_k is not None and top_k < 1:
            raise ValueError(f"top_k must be a positive integer or None, got {top_k}")
        self.top_k = top_k

    def forward(
        self, prop_bounds: Tensor, accessibility: Tensor
    ) -> Tensor:
        """
        Args:
            prop_bounds: ``(|W|, 2)`` or ``(|W|,)`` truth bounds for ϕ.
            accessibility: ``(|W|, |W|)`` accessibility matrix in [0, 1].

        Returns:
            ``(|W|, 2)`` or ``(|W|,)`` truth bounds for □ϕ.
        """
        return F.necessity(
            prop_bounds, accessibility, tau=self.tau.item(), top_k=self.top_k
        )

    def set_tau(self, tau: float) -> None:
        """Set temperature from a float (e.g. for annealing)."""
        t = torch.as_tensor(tau, device=self.tau.device, dtype=self.tau.dtype)
        self.tau.copy_(t)

    def extra_repr(self) -> str:
        return f"tau={self.tau.item():.4f}, top_k={self.top_k}"
forward
forward(prop_bounds: Tensor, accessibility: Tensor) -> Tensor

Parameters:

Name Type Description Default
prop_bounds Tensor

(|W|, 2) or (|W|,) truth bounds for ϕ.

required
accessibility Tensor

(|W|, |W|) accessibility matrix in [0, 1].

required

Returns:

Type Description
Tensor

(|W|, 2) or (|W|,) truth bounds for □ϕ.

Source code in torchmodal/nn/modal.py
def forward(
    self, prop_bounds: Tensor, accessibility: Tensor
) -> Tensor:
    """
    Args:
        prop_bounds: ``(|W|, 2)`` or ``(|W|,)`` truth bounds for ϕ.
        accessibility: ``(|W|, |W|)`` accessibility matrix in [0, 1].

    Returns:
        ``(|W|, 2)`` or ``(|W|,)`` truth bounds for □ϕ.
    """
    return F.necessity(
        prop_bounds, accessibility, tau=self.tau.item(), top_k=self.top_k
    )
set_tau
set_tau(tau: float) -> None

Set temperature from a float (e.g. for annealing).

Source code in torchmodal/nn/modal.py
def set_tau(self, tau: float) -> None:
    """Set temperature from a float (e.g. for annealing)."""
    t = torch.as_tensor(tau, device=self.tau.device, dtype=self.tau.dtype)
    self.tau.copy_(t)

Possibility

Bases: Module

Possibility (Diamond / ♢) neuron.

Implements the differentiable existential quantification over accessible worlds:

.. math:: L_{\Diamond\phi,w} = \operatorname{conv_pool}\tau \bigl{ \tilde{A}{w,w'} + L_{\phi,w'} - 1 \bigr}_{w' \in W}

.. math:: U_{\Diamond\phi,w} = \operatorname{smooth_max}\tau \bigl{ \tilde{A}{w,w'} + U_{\phi,w'} - 1 \bigr}_{w' \in W}

Satisfies modal duality: ♢ϕ ≡ ¬□¬ϕ.

With top_k=k each endpoint aggregates only the k largest of its own conjunction terms (Ã + L - 1 for the lower bound, Ã + U - 1 for the upper), so the true maximum is always kept, the bounds stay sound, and the smooth upper bound is within tau * log(k) of the crisp maximum. See :class:Necessity.

Parameters:

Name Type Description Default
tau float

Temperature for soft aggregation. Default 0.1.

0.1
learnable_tau bool

If True, temperature is learnable. Default False.

False
top_k Optional[int]

If set, aggregate only the top_k largest conjunction terms per world and endpoint. Default None (full row).

None

For temperature annealing during training, update the temperature via :meth:set_tau rather than assigning to .tau.

Example::

>>> diamond = torchmodal.nn.Possibility(tau=0.1)
>>> dia_phi = diamond(prop_bounds, A)
Source code in torchmodal/nn/modal.py
class Possibility(nn.Module):
    r"""Possibility (Diamond / ♢) neuron.

    Implements the differentiable existential quantification over accessible
    worlds:

    .. math::
        L_{\Diamond\phi,w} = \operatorname{conv\_pool}_\tau \bigl\{
            \tilde{A}_{w,w'} + L_{\phi,w'} - 1 \bigr\}_{w' \in W}

    .. math::
        U_{\Diamond\phi,w} = \operatorname{smooth\_max}_\tau \bigl\{
            \tilde{A}_{w,w'} + U_{\phi,w'} - 1 \bigr\}_{w' \in W}

    Satisfies modal duality: ``♢ϕ ≡ ¬□¬ϕ``.

    With ``top_k=k`` each endpoint aggregates only the *k largest* of its
    own conjunction terms (``Ã + L - 1`` for the lower bound, ``Ã + U - 1``
    for the upper), so the true maximum is always kept, the bounds stay
    sound, and the smooth upper bound is within ``tau * log(k)`` of the
    crisp maximum. See :class:`Necessity`.

    Args:
        tau: Temperature for soft aggregation. Default 0.1.
        learnable_tau: If ``True``, temperature is learnable. Default False.
        top_k: If set, aggregate only the ``top_k`` largest conjunction
            terms per world and endpoint. Default ``None`` (full row).

    For temperature annealing during training, update the temperature via
    :meth:`set_tau` rather than assigning to ``.tau``.

    Example::

        >>> diamond = torchmodal.nn.Possibility(tau=0.1)
        >>> dia_phi = diamond(prop_bounds, A)
    """

    def __init__(
        self,
        tau: float = 0.1,
        learnable_tau: bool = False,
        top_k: Optional[int] = None,
    ) -> None:
        super().__init__()
        if learnable_tau:
            self.tau = nn.Parameter(torch.tensor(tau))
        else:
            self.register_buffer("tau", torch.tensor(tau))
        if top_k is not None and top_k < 1:
            raise ValueError(f"top_k must be a positive integer or None, got {top_k}")
        self.top_k = top_k

    def forward(
        self, prop_bounds: Tensor, accessibility: Tensor
    ) -> Tensor:
        """
        Args:
            prop_bounds: ``(|W|, 2)`` or ``(|W|,)`` truth bounds for ϕ.
            accessibility: ``(|W|, |W|)`` accessibility matrix in [0, 1].

        Returns:
            ``(|W|, 2)`` or ``(|W|,)`` truth bounds for ♢ϕ.
        """
        return F.possibility(
            prop_bounds, accessibility, tau=self.tau.item(), top_k=self.top_k
        )

    def set_tau(self, tau: float) -> None:
        """Set temperature from a float (e.g. for annealing)."""
        t = torch.as_tensor(tau, device=self.tau.device, dtype=self.tau.dtype)
        self.tau.copy_(t)

    def extra_repr(self) -> str:
        return f"tau={self.tau.item():.4f}, top_k={self.top_k}"
forward
forward(prop_bounds: Tensor, accessibility: Tensor) -> Tensor

Parameters:

Name Type Description Default
prop_bounds Tensor

(|W|, 2) or (|W|,) truth bounds for ϕ.

required
accessibility Tensor

(|W|, |W|) accessibility matrix in [0, 1].

required

Returns:

Type Description
Tensor

(|W|, 2) or (|W|,) truth bounds for ♢ϕ.

Source code in torchmodal/nn/modal.py
def forward(
    self, prop_bounds: Tensor, accessibility: Tensor
) -> Tensor:
    """
    Args:
        prop_bounds: ``(|W|, 2)`` or ``(|W|,)`` truth bounds for ϕ.
        accessibility: ``(|W|, |W|)`` accessibility matrix in [0, 1].

    Returns:
        ``(|W|, 2)`` or ``(|W|,)`` truth bounds for ♢ϕ.
    """
    return F.possibility(
        prop_bounds, accessibility, tau=self.tau.item(), top_k=self.top_k
    )
set_tau
set_tau(tau: float) -> None

Set temperature from a float (e.g. for annealing).

Source code in torchmodal/nn/modal.py
def set_tau(self, tau: float) -> None:
    """Set temperature from a float (e.g. for annealing)."""
    t = torch.as_tensor(tau, device=self.tau.device, dtype=self.tau.dtype)
    self.tau.copy_(t)

torchmodal.nn.accessibility

torchmodal.nn.accessibility ~~~~~~~~~~~~~~~~~~~~~~~~~~~

Accessibility relation modules for Kripke structures.

Provides four parameterizations:

  • FixedAccessibility: Static, user-defined binary relation.
  • LearnableAccessibility: Direct learnable logit matrix → sigmoid. O(|W|²) parameters — suitable for |W| ≤ ~1000.
  • MetricAccessibility: Metric-learning parameterization using latent embeddings with inner-product kernel. O(d·|W|) parameters — scales to |W| = 20,000+.
  • AttentionAccessibility: Multi-head self-attention over world representations. O(d²) parameters — suitable when worlds have rich feature representations and the accessibility pattern is context-dependent. Addresses the reviewer concern (R1) that the kernel parameterization is not the only sub-quadratic alternative.

Top-k is not an accessibility-module concern. Earlier releases took a top_k argument here and zeroed all but the k largest entries of each row of A before the modal operators saw it. That was unsound: □ and ♢ aggregate (1 - A) + L and A + U - 1, so choosing neighbours by A alone can drop the world whose L / U carries the extremum, and the zeroed entries still enter the log-sum-exp with mass that grows with |W| and drives every bound to [0, 1]. Top-k aggregation now lives on :class:torchmodal.nn.Necessity / :class:~torchmodal.nn.Possibility (top_k=), which select the k extreme aggregation terms per endpoint. top_k here is deprecated and ignored (with a DeprecationWarning).

What remains available here is sparsify=k: a deliberately sparsified relation in which each world accesses only its k most accessible worlds. That is a modelling choice — it defines a different Kripke frame — not an aggregation optimisation, and the operators then reason soundly about the sparsified frame.

FixedAccessibility

Bases: Module

Fixed (non-learnable) accessibility relation.

Wraps a user-defined binary relation matrix as a frozen buffer. Useful for deductive mode where the logical structure is known (e.g., Sudoku constraints, temporal flow, grammatical rules).

Parameters:

Name Type Description Default
relation Tensor

Binary accessibility matrix of shape (|W|, |W|). Values should be 0 or 1.

required
sparsify Optional[int]

If set, keep only the sparsify largest entries of each row and make every other world inaccessible (a different, sparser Kripke frame — a modelling choice, see :func:top_k_mask). Default None.

None
top_k Optional[int]

Deprecated and ignored. Masking A before aggregation was unsound; pass top_k to :class:torchmodal.nn.Necessity / :class:~torchmodal.nn.Possibility instead, or use sparsify for a sparsified relation.

None

Example::

>>> # Sudoku: cells in same row/col/box are accessible
>>> R = build_sudoku_accessibility(9)
>>> access = FixedAccessibility(R)
>>> A = access()  # (81, 81) binary matrix
Source code in torchmodal/nn/accessibility.py
class FixedAccessibility(nn.Module):
    """Fixed (non-learnable) accessibility relation.

    Wraps a user-defined binary relation matrix as a frozen buffer.
    Useful for deductive mode where the logical structure is known
    (e.g., Sudoku constraints, temporal flow, grammatical rules).

    Args:
        relation: Binary accessibility matrix of shape ``(|W|, |W|)``.
            Values should be 0 or 1.
        sparsify: If set, keep only the ``sparsify`` largest entries of
            each row and make every other world inaccessible (a different,
            sparser Kripke frame — a modelling choice, see
            :func:`top_k_mask`). Default ``None``.
        top_k: **Deprecated and ignored.** Masking ``A`` before
            aggregation was unsound; pass ``top_k`` to
            :class:`torchmodal.nn.Necessity` / :class:`~torchmodal.nn.Possibility`
            instead, or use ``sparsify`` for a sparsified relation.

    Example::

        >>> # Sudoku: cells in same row/col/box are accessible
        >>> R = build_sudoku_accessibility(9)
        >>> access = FixedAccessibility(R)
        >>> A = access()  # (81, 81) binary matrix
    """

    #: Declared so mypy knows the registered buffer is a Tensor;
    #: ``nn.Module.__getattr__`` otherwise widens it to
    #: ``Union[Tensor, Module]`` and every use has to be narrowed.
    relation: Tensor

    def __init__(
        self,
        relation: Tensor,
        top_k: Optional[int] = None,
        sparsify: Optional[int] = None,
    ) -> None:
        super().__init__()
        self.register_buffer("relation", relation.float())
        if top_k is not None:
            _warn_top_k_deprecated(self)
        self.sparsify = sparsify

    @property
    def num_worlds(self) -> int:
        return int(self.relation.shape[0])

    def forward(self) -> Tensor:
        """Returns the accessibility matrix ``(|W|, |W|)``."""
        A = self.relation
        if self.sparsify is not None:
            A = top_k_mask(A, self.sparsify)
        return A

    def extra_repr(self) -> str:
        return (
            f"num_worlds={self.num_worlds}, "
            f"sparsify={self.sparsify}"
        )
forward
forward() -> Tensor

Returns the accessibility matrix (|W|, |W|).

Source code in torchmodal/nn/accessibility.py
def forward(self) -> Tensor:
    """Returns the accessibility matrix ``(|W|, |W|)``."""
    A = self.relation
    if self.sparsify is not None:
        A = top_k_mask(A, self.sparsify)
    return A

LearnableAccessibility

Bases: Module

Learnable accessibility relation via direct logit matrix.

Parameterizes R as a matrix of learnable logits passed through sigmoid: A = σ(logits). Suitable for small-to-medium world sets (|W| ≤ ~1000).

The parameter space is O(|W|²).

Parameters:

Name Type Description Default
num_worlds int

Number of possible worlds |W|.

required
init_bias float

Initial bias for logits. Negative values encode a "prior of distrust" (default -2.0).

-2.0
reflexive bool

If True, enforce self-accessibility (diagonal = 1). Default True.

True
sparsify Optional[int]

If set, keep only the sparsify largest entries of each row after the sigmoid and make every other world inaccessible (a sparser Kripke frame — a modelling choice, see :func:top_k_mask). Default None.

None
top_k Optional[int]

Deprecated and ignored. Masking A before aggregation was unsound; pass top_k to :class:torchmodal.nn.Necessity / :class:~torchmodal.nn.Possibility instead, or use sparsify for a sparsified relation.

None

Example::

>>> access = LearnableAccessibility(7, reflexive=True)
>>> A = access()  # (7, 7) matrix in [0, 1]
Source code in torchmodal/nn/accessibility.py
class LearnableAccessibility(nn.Module):
    """Learnable accessibility relation via direct logit matrix.

    Parameterizes R as a matrix of learnable logits passed through
    sigmoid: ``A = σ(logits)``. Suitable for small-to-medium world
    sets (|W| ≤ ~1000).

    The parameter space is O(|W|²).

    Args:
        num_worlds: Number of possible worlds |W|.
        init_bias: Initial bias for logits. Negative values encode a
            "prior of distrust" (default -2.0).
        reflexive: If ``True``, enforce self-accessibility (diagonal = 1).
            Default ``True``.
        sparsify: If set, keep only the ``sparsify`` largest entries of
            each row after the sigmoid and make every other world
            inaccessible (a sparser Kripke frame — a modelling choice, see
            :func:`top_k_mask`). Default ``None``.
        top_k: **Deprecated and ignored.** Masking ``A`` before
            aggregation was unsound; pass ``top_k`` to
            :class:`torchmodal.nn.Necessity` / :class:`~torchmodal.nn.Possibility`
            instead, or use ``sparsify`` for a sparsified relation.

    Example::

        >>> access = LearnableAccessibility(7, reflexive=True)
        >>> A = access()  # (7, 7) matrix in [0, 1]
    """

    def __init__(
        self,
        num_worlds: int,
        init_bias: float = -2.0,
        reflexive: bool = True,
        top_k: Optional[int] = None,
        sparsify: Optional[int] = None,
    ) -> None:
        super().__init__()
        self._num_worlds = num_worlds
        self.reflexive = reflexive
        if top_k is not None:
            _warn_top_k_deprecated(self)
        self.sparsify = sparsify

        self.logits = nn.Parameter(
            torch.full((num_worlds, num_worlds), init_bias)
        )

        if reflexive:
            # Initialize diagonal to high logit (self-trust)
            with torch.no_grad():
                self.logits.diagonal().fill_(5.0)

    @property
    def num_worlds(self) -> int:
        return self._num_worlds

    def forward(self) -> Tensor:
        """Returns the accessibility matrix ``(|W|, |W|)`` in [0, 1]."""
        A = torch.sigmoid(self.logits)

        if self.reflexive:
            # Clamp diagonal to 1.0
            A = A.clone()
            A.fill_diagonal_(1.0)

        if self.sparsify is not None:
            A = top_k_mask(A, self.sparsify)

        return A

    def extra_repr(self) -> str:
        return (
            f"num_worlds={self._num_worlds}, "
            f"reflexive={self.reflexive}, "
            f"sparsify={self.sparsify}"
        )
forward
forward() -> Tensor

Returns the accessibility matrix (|W|, |W|) in [0, 1].

Source code in torchmodal/nn/accessibility.py
def forward(self) -> Tensor:
    """Returns the accessibility matrix ``(|W|, |W|)`` in [0, 1]."""
    A = torch.sigmoid(self.logits)

    if self.reflexive:
        # Clamp diagonal to 1.0
        A = A.clone()
        A.fill_diagonal_(1.0)

    if self.sparsify is not None:
        A = top_k_mask(A, self.sparsify)

    return A

MetricAccessibility

Bases: Module

Scalable metric-learning accessibility relation.

Maps each world to a latent embedding and computes accessibility via a kernel function:

.. math:: A(w_i, w_j) = \sigma\bigl(h_{w_i}^\top h_{w_j}\bigr)

This reduces the parameter space from O(|W|²) to O(d·|W|) and enables scaling to |W| = 20,000+ on a single GPU.

The encoder can optionally accept external features per world.

Parameters:

Name Type Description Default
num_worlds int

Number of possible worlds |W|.

required
embed_dim int

Embedding dimension d. Default 64.

64
input_dim Optional[int]

If provided, the encoder takes external features of this dimension. Otherwise, uses learnable embeddings.

None
hidden_dim int

Hidden dimension of the encoder MLP. Default 128.

128
reflexive bool

Enforce self-accessibility. Default True.

True
sparsify Optional[int]

If set, keep only the sparsify largest entries of each row and make every other world inaccessible (a sparser Kripke frame — a modelling choice, see :func:top_k_mask). Default None.

None
top_k Optional[int]

Deprecated and ignored. Masking A before aggregation was unsound; pass top_k to :class:torchmodal.nn.Necessity / :class:~torchmodal.nn.Possibility instead, or use sparsify for a sparsified relation.

None

Example::

>>> access = MetricAccessibility(1000, embed_dim=64)
>>> A = access()  # (1000, 1000) accessibility matrix
>>> # With external features:
>>> access = MetricAccessibility(100, embed_dim=32, input_dim=384)
>>> A = access(features)  # features: (100, 384)
Source code in torchmodal/nn/accessibility.py
class MetricAccessibility(nn.Module):
    """Scalable metric-learning accessibility relation.

    Maps each world to a latent embedding and computes accessibility
    via a kernel function:

    .. math::
        A(w_i, w_j) = \\sigma\\bigl(h_{w_i}^\\top h_{w_j}\\bigr)

    This reduces the parameter space from O(|W|²) to O(d·|W|) and
    enables scaling to |W| = 20,000+ on a single GPU.

    The encoder can optionally accept external features per world.

    Args:
        num_worlds: Number of possible worlds |W|.
        embed_dim: Embedding dimension *d*. Default 64.
        input_dim: If provided, the encoder takes external features of
            this dimension. Otherwise, uses learnable embeddings.
        hidden_dim: Hidden dimension of the encoder MLP. Default 128.
        reflexive: Enforce self-accessibility. Default ``True``.
        sparsify: If set, keep only the ``sparsify`` largest entries of
            each row and make every other world inaccessible (a sparser
            Kripke frame — a modelling choice, see :func:`top_k_mask`).
            Default ``None``.
        top_k: **Deprecated and ignored.** Masking ``A`` before
            aggregation was unsound; pass ``top_k`` to
            :class:`torchmodal.nn.Necessity` / :class:`~torchmodal.nn.Possibility`
            instead, or use ``sparsify`` for a sparsified relation.

    Example::

        >>> access = MetricAccessibility(1000, embed_dim=64)
        >>> A = access()  # (1000, 1000) accessibility matrix
        >>> # With external features:
        >>> access = MetricAccessibility(100, embed_dim=32, input_dim=384)
        >>> A = access(features)  # features: (100, 384)
    """

    def __init__(
        self,
        num_worlds: int,
        embed_dim: int = 64,
        input_dim: Optional[int] = None,
        hidden_dim: int = 128,
        reflexive: bool = True,
        top_k: Optional[int] = None,
        sparsify: Optional[int] = None,
    ) -> None:
        super().__init__()
        self._num_worlds = num_worlds
        self.embed_dim = embed_dim
        self.reflexive = reflexive
        if top_k is not None:
            _warn_top_k_deprecated(self)
        self.sparsify = sparsify

        # Exactly one of these is populated; annotate before the branch so
        # each is declared once.
        self.encoder: Optional[nn.Sequential]
        self.embeddings: Optional[nn.Parameter]
        if input_dim is not None:
            # Encoder from external features
            self.encoder = nn.Sequential(
                nn.Linear(input_dim, hidden_dim),
                nn.ReLU(),
                nn.Linear(hidden_dim, embed_dim),
            )
            self.embeddings = None
        else:
            # Learnable embeddings per world
            self.encoder = None
            self.embeddings = nn.Parameter(
                torch.randn(num_worlds, embed_dim) * 0.01
            )

    @property
    def num_worlds(self) -> int:
        return self._num_worlds

    def forward(self, features: Optional[Tensor] = None) -> Tensor:
        """Compute the accessibility matrix.

        Args:
            features: Optional external features ``(|W|, input_dim)``.
                Required if ``input_dim`` was set at construction.

        Returns:
            Accessibility matrix ``(|W|, |W|)`` in [0, 1].
        """
        if self.encoder is not None:
            if features is None:
                raise ValueError(
                    "MetricAccessibility with input_dim requires features"
                )
            h = self.encoder(features)
        else:
            h = self.embeddings

        # Kernel: inner product → sigmoid
        A = torch.sigmoid(h @ h.t())

        if self.reflexive:
            A = A.clone()
            A.fill_diagonal_(1.0)

        if self.sparsify is not None:
            A = top_k_mask(A, self.sparsify)

        return A

    def extra_repr(self) -> str:
        return (
            f"num_worlds={self._num_worlds}, "
            f"embed_dim={self.embed_dim}, "
            f"reflexive={self.reflexive}, "
            f"sparsify={self.sparsify}"
        )
forward
forward(features: Optional[Tensor] = None) -> Tensor

Compute the accessibility matrix.

Parameters:

Name Type Description Default
features Optional[Tensor]

Optional external features (|W|, input_dim). Required if input_dim was set at construction.

None

Returns:

Type Description
Tensor

Accessibility matrix (|W|, |W|) in [0, 1].

Source code in torchmodal/nn/accessibility.py
def forward(self, features: Optional[Tensor] = None) -> Tensor:
    """Compute the accessibility matrix.

    Args:
        features: Optional external features ``(|W|, input_dim)``.
            Required if ``input_dim`` was set at construction.

    Returns:
        Accessibility matrix ``(|W|, |W|)`` in [0, 1].
    """
    if self.encoder is not None:
        if features is None:
            raise ValueError(
                "MetricAccessibility with input_dim requires features"
            )
        h = self.encoder(features)
    else:
        h = self.embeddings

    # Kernel: inner product → sigmoid
    A = torch.sigmoid(h @ h.t())

    if self.reflexive:
        A = A.clone()
        A.fill_diagonal_(1.0)

    if self.sparsify is not None:
        A = top_k_mask(A, self.sparsify)

    return A

AttentionAccessibility

Bases: Module

Attention-based accessibility relation.

Uses multi-head self-attention over world representations to compute a context-dependent accessibility matrix. Unlike :class:MetricAccessibility (which uses a fixed inner-product kernel), attention weights are input-dependent and can capture asymmetric relationships naturally.

The parameter count is O(d²) — independent of |W| — making this suitable for settings where worlds have rich feature representations (e.g., sentence embeddings in the Diplomacy experiment).

This addresses Reviewer 1's observation that "if worlds were a space of rich state representations rather than indices, directly learning a kernel does not require quadratic parameters" by providing an alternative that operates entirely in feature space.

Parameters:

Name Type Description Default
input_dim int

Dimension of per-world feature vectors.

required
num_heads int

Number of attention heads. Default 4.

4
reflexive bool

Enforce self-accessibility. Default True.

True
sparsify Optional[int]

If set, keep only the sparsify largest entries of each row and make every other world inaccessible (a sparser Kripke frame — a modelling choice, see :func:top_k_mask). Default None.

None
top_k Optional[int]

Deprecated and ignored. Masking A before aggregation was unsound; pass top_k to :class:torchmodal.nn.Necessity / :class:~torchmodal.nn.Possibility instead, or use sparsify for a sparsified relation.

None

Example::

>>> access = AttentionAccessibility(input_dim=384, num_heads=4)
>>> features = torch.randn(7, 384)  # 7 worlds, 384-d features
>>> A = access(features)  # (7, 7) accessibility matrix
Source code in torchmodal/nn/accessibility.py
class AttentionAccessibility(nn.Module):
    """Attention-based accessibility relation.

    Uses multi-head self-attention over world representations to compute
    a context-dependent accessibility matrix.  Unlike
    :class:`MetricAccessibility` (which uses a fixed inner-product
    kernel), attention weights are input-dependent and can capture
    asymmetric relationships naturally.

    The parameter count is O(d²) — independent of |W| — making this
    suitable for settings where worlds have rich feature representations
    (e.g., sentence embeddings in the Diplomacy experiment).

    This addresses Reviewer 1's observation that "if worlds were a space
    of rich state representations rather than indices, directly learning
    a kernel does not require quadratic parameters" by providing an
    alternative that operates entirely in feature space.

    Args:
        input_dim: Dimension of per-world feature vectors.
        num_heads: Number of attention heads. Default 4.
        reflexive: Enforce self-accessibility. Default ``True``.
        sparsify: If set, keep only the ``sparsify`` largest entries of
            each row and make every other world inaccessible (a sparser
            Kripke frame — a modelling choice, see :func:`top_k_mask`).
            Default ``None``.
        top_k: **Deprecated and ignored.** Masking ``A`` before
            aggregation was unsound; pass ``top_k`` to
            :class:`torchmodal.nn.Necessity` / :class:`~torchmodal.nn.Possibility`
            instead, or use ``sparsify`` for a sparsified relation.

    Example::

        >>> access = AttentionAccessibility(input_dim=384, num_heads=4)
        >>> features = torch.randn(7, 384)  # 7 worlds, 384-d features
        >>> A = access(features)  # (7, 7) accessibility matrix
    """

    def __init__(
        self,
        input_dim: int,
        num_heads: int = 4,
        reflexive: bool = True,
        top_k: Optional[int] = None,
        sparsify: Optional[int] = None,
    ) -> None:
        super().__init__()
        self.input_dim = input_dim
        self.num_heads = num_heads
        self.reflexive = reflexive
        if top_k is not None:
            _warn_top_k_deprecated(self)
        self.sparsify = sparsify

        self.attn = nn.MultiheadAttention(
            embed_dim=input_dim,
            num_heads=num_heads,
            batch_first=True,
        )
        self.proj = nn.Linear(input_dim, 1)

    def forward(self, features: Tensor) -> Tensor:
        """Compute the accessibility matrix from world features.

        Args:
            features: Per-world features ``(|W|, input_dim)``.

        Returns:
            Accessibility matrix ``(|W|, |W|)`` in [0, 1].
        """
        # (1, |W|, d) for batch-first MHA
        x = features.unsqueeze(0)
        attn_out, attn_weights = self.attn(x, x, x)
        # attn_weights: (1, |W|, |W|) — already in [0, 1] (softmax)
        A = attn_weights.squeeze(0)

        if self.reflexive:
            A = A.clone()
            A.fill_diagonal_(1.0)

        if self.sparsify is not None:
            A = top_k_mask(A, self.sparsify)

        return cast(Tensor, A)

    def extra_repr(self) -> str:
        return (
            f"input_dim={self.input_dim}, "
            f"num_heads={self.num_heads}, "
            f"reflexive={self.reflexive}, "
            f"sparsify={self.sparsify}"
        )
forward
forward(features: Tensor) -> Tensor

Compute the accessibility matrix from world features.

Parameters:

Name Type Description Default
features Tensor

Per-world features (|W|, input_dim).

required

Returns:

Type Description
Tensor

Accessibility matrix (|W|, |W|) in [0, 1].

Source code in torchmodal/nn/accessibility.py
def forward(self, features: Tensor) -> Tensor:
    """Compute the accessibility matrix from world features.

    Args:
        features: Per-world features ``(|W|, input_dim)``.

    Returns:
        Accessibility matrix ``(|W|, |W|)`` in [0, 1].
    """
    # (1, |W|, d) for batch-first MHA
    x = features.unsqueeze(0)
    attn_out, attn_weights = self.attn(x, x, x)
    # attn_weights: (1, |W|, |W|) — already in [0, 1] (softmax)
    A = attn_weights.squeeze(0)

    if self.reflexive:
        A = A.clone()
        A.fill_diagonal_(1.0)

    if self.sparsify is not None:
        A = top_k_mask(A, self.sparsify)

    return cast(Tensor, A)

top_k_mask

top_k_mask(A: Tensor, k: int) -> Tensor

Sparsify an accessibility matrix to its k largest entries per row.

For each world (row), only the k highest accessibility values are kept; all others are set to 0, i.e. those worlds become inaccessible. This defines a different (sparser) Kripke frame and is what the sparsify= option of the accessibility modules applies.

.. warning:: This is a modelling choice, not an aggregation optimisation. It does not reduce the cost of □ / ♢ (the operators still aggregate over the full (|W|, |W|) row) and it must not be used to emulate top-k aggregation: neighbours are chosen by A alone rather than by the aggregated terms, and the zeroed entries still enter the log-sum-exp with term 1 + L each, so the bounds drift to [0, 1] as |W| grows. For sound top-k aggregation with a tau * log(k) gap use top_k= on :func:torchmodal.functional.necessity / :func:~torchmodal.functional.possibility or the :class:torchmodal.nn.Necessity / :class:~torchmodal.nn.Possibility modules.

Parameters:

Name Type Description Default
A Tensor

Accessibility matrix of shape (|W|, |W|).

required
k int

Number of neighbors to retain per world.

required

Returns:

Type Description
Tensor

Sparsified accessibility matrix of the same shape.

Source code in torchmodal/nn/accessibility.py
def top_k_mask(A: Tensor, k: int) -> Tensor:
    """Sparsify an accessibility matrix to its *k* largest entries per row.

    For each world (row), only the *k* highest accessibility values are
    kept; all others are set to 0, i.e. those worlds become
    *inaccessible*. This defines a different (sparser) Kripke frame and is
    what the ``sparsify=`` option of the accessibility modules applies.

    .. warning::
       This is a **modelling choice, not an aggregation optimisation**.
       It does not reduce the cost of □ / ♢ (the operators still aggregate
       over the full ``(|W|, |W|)`` row) and it must not be used to emulate
       top-k aggregation: neighbours are chosen by ``A`` alone rather than
       by the aggregated terms, and the zeroed entries still enter the
       log-sum-exp with term ``1 + L`` each, so the bounds drift to
       ``[0, 1]`` as ``|W|`` grows. For sound top-k aggregation with a
       ``tau * log(k)`` gap use ``top_k=`` on
       :func:`torchmodal.functional.necessity` /
       :func:`~torchmodal.functional.possibility` or the
       :class:`torchmodal.nn.Necessity` / :class:`~torchmodal.nn.Possibility`
       modules.

    Args:
        A: Accessibility matrix of shape ``(|W|, |W|)``.
        k: Number of neighbors to retain per world.

    Returns:
        Sparsified accessibility matrix of the same shape.
    """
    if k >= A.shape[-1]:
        return A
    topk_vals, _ = torch.topk(A, k, dim=-1)
    threshold = topk_vals[..., -1:]
    mask = (A >= threshold).float()
    return A * mask