MCD: d-separation for causal self-explanation#
Liu, Wang, Wang, Li, Deng, Zhang and Qiu, 2023, D-Separation for Causal Self-Explanation, NeurIPS 2023. Reference implementation: jugechengzi/Rationalization-MCD.
Problem#
Asking the highlight to predict the label rewards any subset that carries the label, and a spurious feature carries it as well as a causal one. The usual answer is a penalty per known spurious pattern, which is a list somebody has to write and a corpus is free to fall outside of. MCD asks for a property of the whole document instead. If the highlight carries what determines the label, the rest of the document adds nothing once the highlight is known. That is conditional independence between the label and the remainder, given the highlight.
Method#
A second prediction is made from the full input, and the highlight is asked to make the two predictions agree. Given a highlight that d-separates the label from the rest of the input, the two predictions have the same information behind them. The divergence between them is therefore the signal the generator is trained on, and no list of spurious patterns is needed.
%%{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 --> PH["predictor<br/>reads the highlight"]
X --> PF["predictor<br/>reads every word"]
PH --> D{{"discrepancy<br/>the generator minimises"}}
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,PH,PF head
class X,H value
class D term
Training alternates two phases. The predictor phase trains the predictor on a selection it is handed and may not move. The generator phase trains the generator while the predictor stays fixed, so neither module can reduce the divergence by adapting to the other.
The predictor learns to classify from the highlight and from the whole document, the generator moves the highlight until those two answers agree, and the two penalties on the mask are scored in both phases.
Training#
Two optimizers and two forward passes per batch, driven by the model rather than by Lightning.
The predictor phase selects, detaches the selection, and scores the classification terms over the highlight and over the full input.
Both optimizers step, since the penalties on the mask are shared and reach the generator here as well.
The generator phase selects again and scores the discrepancy term.
The generator’s optimizer steps alone.
Two properties follow the reference implementation.
The shared criteria bind to the selection the generator produced rather than to the detached copy.
The generator therefore takes two steps per batch on them and one on its own term, as train_util.train_decouple_causal2 does upstream.
Each phase also draws its own selection, since phase_forward selects once per phase.
The two phases of a batch therefore optimise different masks of it.
Implementation#
MCD is one of the two models built on PhasedSPP, which owns the phases, the three loss lists and the manual training step.
What MCD adds is two methods naming the passes its criteria bind to.
class MCD(PhasedSPP):
"""Rationalizer trained against selected-input and full-input predictions.
A predictor reading the full input guides the generator: a highlight that
d-separates the label from the rest of the input makes the selected-input
and full-input predictions agree.
Liu, Wang, Wang, Li, Deng, Zhang and Qiu, 2023, *D-Separation for Causal
Self-Explanation*, NeurIPS 2023.
Reference implementation:
<https://github.com/jugechengzi/Rationalization-MCD>.
The phases and the three loss lists are
:class:`~pyhighlights.components.models.spp.phased.PhasedSPP`'s. What MCD
adds is the pass they are scored over: the predictor reads the highlight,
and every namespace exposes ``full_class_logits`` beside the
selected-input fields.
"""
def phase_class_logits(
self, input_data: InputData, highlight_mask: th.Tensor, selection: th.Tensor
) -> th.Tensor:
"""The highlight pass, which MCD trains on: in the graph, per phase."""
return self.predict(input_data, selection)
def extra_logits(
self, input_data: InputData, selection: th.Tensor
) -> Dict[str, th.Tensor]:
"""The full input, which is what the highlight has to agree with."""
return {"full_class_logits": self.predict_full(input_data)}
What |
Where |
|---|---|
The alternation itself |
|
Generator stepped alone in the generator phase |
|
Detached selection for the predictor phase |
|
The highlight pass |
|
The full-input pass |
|
Highlight supervision is refused rather than ignored, since a supervision loss appended to the flat list would be dropped before the first batch.
It has to name the phase it belongs to and go in shared_losses.
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 |
|---|---|
|
|
|
|
Three loss lists rather than one, and each names a phase.
Field |
Default |
|---|---|
|
sparsity and contiguity, scored in both phases |
|
classification over the highlight, and classification over the full input |
|
the discrepancy between those two predictions |
from pyhighlights.configurations.keys import GRU_MCD, TOY_TASK
Registry.from_key(TOY_TASK, model=GRU_MCD, save_path="results", seeds=[0, 1])
predictor_backbone is required, since the highlight and the full input are read by an encoder of the predictor’s own.
API#
- class pyhighlights.components.models.spp.mcd.MCD(selector_backbones, selectors, predictor, predictor_backbone, shared_losses, predictor_losses, generator_losses, **kwargs)[source]#
Bases:
PhasedSPPRationalizer trained against selected-input and full-input predictions.
A predictor reading the full input guides the generator: a highlight that d-separates the label from the rest of the input makes the selected-input and full-input predictions agree.
Liu, Wang, Wang, Li, Deng, Zhang and Qiu, 2023, D-Separation for Causal Self-Explanation, NeurIPS 2023. Reference implementation: <jugechengzi/Rationalization-MCD>.
The phases and the three loss lists are
PhasedSPP’s. What MCD adds is the pass they are scored over: the predictor reads the highlight, and every namespace exposesfull_class_logitsbeside the selected-input 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]])
- class pyhighlights.components.models.spp.phased.PhasedSPP(selector_backbones, selectors, predictor, predictor_backbone, shared_losses, predictor_losses, generator_losses, **kwargs)[source]#
Bases:
SPPA rationalizer whose criteria belong to one of two training phases.
Two architectures are built this way. The predictor phase trains the predictor on a selection it is handed and may not move. The generator phase steps the generator alone, so the predictor stays fixed. What differs between them is which criteria go in which list and what the predictor is asked to read, and that is what the subclasses say.
shared_lossesare scored in both phases,predictor_lossesin the predictor phase andgenerator_lossesin the generator phase. Evaluation reports every group, since a validation number is about the model rather than about a phase of its training.Two properties of the loop read oddly until they are checked against the implementations it reproduces, so both are stated here.
The shared criteria bind to the selection the generator produced rather than to the detached copy the predictor reads. They therefore reach the generator in the predictor phase as well as in its own, and the generator’s optimizer is stepped in both. The generator takes two steps per batch on the shared criteria and one on the phase-specific term. Both reference implementations do the same:
train_util.train_decouple_causal2of <jugechengzi/Rationalization-MCD> adds the sparsity and continuity terms to its classification loss and stepsopt_genbesideopt_pred, andtrain_util.train_adv_causalof <jugechengzi/Rationalization-MRD> does the same under--gen_sparse, which defaults to 1 and is what its README runs. A study that wants the other arrangement registersshared_lossesempty, which is what turning that flag off amounts to.And each phase draws its own selection:
phase_forwardselects once per phase, so the two phases of a batch optimize different masks of it. That is the references again, which callget_rationalein each phase over agumbel_softmaxwithhard=True.The training
lossis therefore not the validationloss. Training logs the sum of the two phases, which counts the shared criteria twice and scores each phase on its own selection. Validation scores every criterion once, on one selection.Under data parallelism, every parameter has to receive a gradient in each phase. With
shared_lossesempty, the generator receives none in the predictor phase. Such a model needsstrategy="ddp_find_unused_parameters_true".- 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]])
- abstractmethod extra_logits(input_data, selection)[source]#
The predictor passes this model’s criteria bind to, beside the highlight’s own fields.
- Return type:
Dict[str,Tensor]- Parameters:
input_data (InputData)
selection (Tensor)
- abstractmethod phase_class_logits(input_data, highlight_mask, selection)[source]#
What the predictor makes of the highlight, however the model asks.
highlight_maskis what the generator produced andselectionis the copy this phase hands the predictor, detached in the predictor phase. A model that trains on the highlight reads the second; one that trains on the complement reads neither and takes this pass outside the graph.- Return type:
Tensor- Parameters:
input_data (InputData)
highlight_mask (Tensor)
selection (Tensor)
MCD: selected-input and full-input predictions trained in two phases.