MRD: maximizing the remaining discrepancy#

Liu, Deng, Niu, Wang, Wang, Zhang and Li, 2024, Is the MMI Criterion Necessary for Interpretability? Degenerating Non-causal Features to Plain Noise for Self-Rationalization, NeurIPS 2024. Reference implementation: jugechengzi/Rationalization-MRD.

Problem#

Maximum mutual information asks the highlight to predict the label. A spurious feature correlated with the label answers that question as well as a causal one. Every architecture up to this point attacks the cooperation between the two modules, and none of them changes the question being asked. MRD changes the question.

Method#

Instead of asking what the highlight can say, MRD asks what is left once the highlight is removed. Removing plain noise leaves the conditional distribution of the remainder unchanged, and so does removing a spurious feature. Only removing the causal features moves it. A corpus full of spurious features therefore behaves like a clean one, and no penalty for each pattern is needed. The generator therefore maximises the divergence between what the complement predicts and what the whole document predicts.

        %%{init: {"theme": "base", "themeVariables": {"fontSize": "17px", "fontFamily": "Lato, sans-serif", "lineColor": "#37474f", "primaryTextColor": "#102027", "edgeLabelBackground": "#ffffff"}, "flowchart": {"nodeSpacing": 55, "rankSpacing": 70, "padding": 14, "curve": "basis"}}%%
flowchart LR
    X["input x"] --> SEL["generator"] --> H["highlight h"]
    H -. "removed from x" .-> C["complement"]
    C --> PC["predictor<br/>reads the complement"]
    X --> PF["predictor<br/>reads every word"]
    PC --> D{{"discrepancy<br/>the generator maximises"}}
    PF --> D

    classDef shared fill:#cfe3ff,stroke:#1a4f9c,stroke-width:2px,color:#0b2545
    classDef head fill:#ffffff,stroke:#37474f,stroke-width:2px,color:#102027
    classDef value fill:#d7f0dc,stroke:#1e6b34,stroke-width:2px,color:#0d3018
    classDef term fill:#ffe9c7,stroke:#a35c00,stroke-width:2px,color:#3d2100
    class SEL,PC,PF head
    class X,H,C value
    class D term
    
\[\begin{split}\begin{aligned} \text{predictor phase:} \quad & \min_{\phi} \;\; \mathcal{L}_{\text{cls}}\big(p_\phi(y \mid (1 - h) \odot x),\, y\big) + \mathcal{L}_{\text{cls}}\big(p_\phi(y \mid x),\, y\big) \\ \text{generator phase:} \quad & \max_{\theta} \;\; \mathcal{D}\big(p_\phi(y \mid (1 - h) \odot x) \,\|\, p_\phi(y \mid x)\big) - \lambda_s \Omega_s(h) - \lambda_c \Omega_c(h) \end{aligned}\end{split}\]

Two consequences set this model apart from the others. The predictor never trains on the highlight, since it is trained on the complement and on the full input. The highlight pass exists only so the metrics have something to score. The generator maximises a divergence rather than minimising one, which in the library is a loss registered with a negative coefficient.

Training#

The phases are MCD’s, and only what each phase reads differs.

  1. The predictor phase selects, detaches the selection, and trains the predictor on the complement and on the full input.

  2. Both optimizers step, since the penalties on the mask are shared.

  3. The generator phase selects again and scores the remaining-discrepancy term.

  4. The generator’s optimizer steps alone.

  5. The highlight pass runs under no_grad, since nothing trains on it and it exists to be measured.

Implementation#

Like MCD, MRD is a PhasedSPP that names the passes its criteria bind to.

class MRD(PhasedSPP):
    """Rationalizer trained by what is left once the highlight is removed.

    Maximum mutual information asks the highlight to predict the label, which
    a spurious feature correlated with the label answers just as well. MRD
    asks the opposite question: remove the highlight, and what remains should
    stop looking like the whole input. Removing plain noise or a spurious
    feature leaves the conditional distribution of the rest unchanged, so only
    the causal features move it. A corpus full of spurious features therefore
    behaves like a clean one, and no penalty per spurious feature is needed.

    Two consequences set this model apart from the others. The predictor
    never trains on the highlight: it is trained on the
    **complement** and on the full input, and the highlight pass exists only
    so the metrics have something to score. And the generator maximizes a
    divergence rather than minimizing one, which is a loss with a negative
    coefficient.

    Liu, Deng, Niu, Wang, Wang, Zhang and Li, 2024, *Is the MMI Criterion
    Necessary for Interpretability? Degenerating Non-causal Features to Plain
    Noise for Self-Rationalization*, NeurIPS 2024.
    Paper: <https://proceedings.neurips.cc/paper_files/paper/2024/hash/d53d51e88d92d3723755f6d425bc513b-Abstract-Conference.html>.
    Reference implementation:
    <https://github.com/jugechengzi/Rationalization-MRD>.

    The phases and the three loss lists are
    :class:`~pyhighlights.components.models.spp.phased.PhasedSPP`'s, as MCD's
    are. Every namespace carries ``complement_class_logits`` and
    ``full_class_logits`` beside the highlight's own fields.
    """

    def phase_class_logits(
        self, input_data: InputData, highlight_mask: th.Tensor, selection: th.Tensor
    ) -> th.Tensor:
        """The highlight pass, outside the graph.

        Nothing trains on it, and the reference implementation keeps it for
        the same reason this does: a number to score the model by, not a term.
        """
        with th.no_grad():
            return self.predict(input_data, highlight_mask.detach())

    def extra_logits(
        self, input_data: InputData, selection: th.Tensor
    ) -> Dict[str, th.Tensor]:
        """The two passes MRD trains on: the complement, and the full input."""
        return {
            "complement_class_logits": self.predict_complement(input_data, selection),
            "full_class_logits": self.predict_full(input_data),
        }

What

Where

The highlight pass, outside the graph

MRD.phase_class_logits

The complement and full-input passes

MRD.extra_logits

The complement itself

SPP.predict_complement

The alternation

PhasedSPP.training_step

A highlight covering every valid word leaves the complement empty, which the backbones pool to zeros. That measures a model that kept everything, so nothing guards it.

Differences from the reference implementation#

The sparsity and contiguity penalties also differ from the reference implementation, as Differences from the reference implementation on the FR page describes.

Configuration#

Key

Configuration

GRU_MRD

GRUMRDConfig

TRANSFORMER_MRD

TransformerMRDConfig

Field

Default

shared_losses

sparsity and contiguity, scored in both phases

predictor_losses

classification over the complement, and classification over the full input

generator_losses

the remaining-discrepancy term

from pyhighlights.configurations.keys import GRU_MRD, TOY_TASK

Registry.from_key(TOY_TASK, model=GRU_MRD, save_path="results", seeds=[0, 1])

API#

class pyhighlights.components.models.spp.mrd.MRD(selector_backbones, selectors, predictor, predictor_backbone, shared_losses, predictor_losses, generator_losses, **kwargs)[source]#

Bases: PhasedSPP

Rationalizer trained by what is left once the highlight is removed.

Maximum mutual information asks the highlight to predict the label, which a spurious feature correlated with the label answers just as well. MRD asks the opposite question: remove the highlight, and what remains should stop looking like the whole input. Removing plain noise or a spurious feature leaves the conditional distribution of the rest unchanged, so only the causal features move it. A corpus full of spurious features therefore behaves like a clean one, and no penalty per spurious feature is needed.

Two consequences set this model apart from the others. The predictor never trains on the highlight: it is trained on the complement and on the full input, and the highlight pass exists only so the metrics have something to score. And the generator maximizes a divergence rather than minimizing one, which is a loss with a negative coefficient.

Liu, Deng, Niu, Wang, Wang, Zhang and Li, 2024, Is the MMI Criterion Necessary for Interpretability? Degenerating Non-causal Features to Plain Noise for Self-Rationalization, NeurIPS 2024. Paper: <https://proceedings.neurips.cc/paper_files/paper/2024/hash/d53d51e88d92d3723755f6d425bc513b-Abstract-Conference.html>. Reference implementation: <jugechengzi/Rationalization-MRD>.

The phases and the three loss lists are PhasedSPP’s, as MCD’s are. Every namespace carries complement_class_logits and full_class_logits beside the highlight’s own fields.

Parameters:
  • selector_backbones (RegistrationKey[SPPBackbone])

  • selectors (RegistrationKey[SPPSelector])

  • predictor (RegistrationKey[SPPPredictor])

  • predictor_backbone (RegistrationKey[SPPBackbone] | None)

  • shared_losses (List[RegistrationKey[Loss]])

  • predictor_losses (List[RegistrationKey[Loss]])

  • generator_losses (List[RegistrationKey[Loss]])

extra_logits(input_data, selection)[source]#

The two passes MRD trains on: the complement, and the full input.

Return type:

Dict[str, Tensor]

Parameters:
  • input_data (InputData)

  • selection (Tensor)

phase_class_logits(input_data, highlight_mask, selection)[source]#

The highlight pass, outside the graph.

Nothing trains on it, and the reference implementation keeps it for the same reason this does: a number to score the model by, not a term.

Return type:

Tensor

Parameters:
  • input_data (InputData)

  • highlight_mask (Tensor)

  • selection (Tensor)

MRD: a generator trained to make the complement stop predicting the label.