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
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.
The predictor phase selects, detaches the selection, and trains the predictor on the complement and on the full input.
Both optimizers step, since the penalties on the mask are shared.
The generator phase selects again and scores the remaining-discrepancy term.
The generator’s optimizer steps alone.
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 |
|
The complement and full-input passes |
|
The complement itself |
|
The alternation |
|
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 |
|---|---|
|
|
|
|
Field |
Default |
|---|---|
|
sparsity and contiguity, scored in both phases |
|
classification over the complement, and classification over the full input |
|
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:
PhasedSPPRationalizer 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 carriescomplement_class_logitsandfull_class_logitsbeside 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.