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.

\[\begin{split}\begin{aligned} \text{predictor phase:} \quad & \min_{\phi} \;\; \mathcal{L}_{\text{cls}}\big(p_\phi(y \mid h \odot x),\, y\big) + \mathcal{L}_{\text{cls}}\big(p_\phi(y \mid x),\, y\big) \\ \text{generator phase:} \quad & \min_{\theta} \;\; \mathcal{D}\big(p_\phi(y \mid h \odot x) \,\|\, p_\phi(y \mid x)\big) + \lambda_s \Omega_s(h) + \lambda_c \Omega_c(h) \end{aligned}\end{split}\]

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.

  1. The predictor phase selects, detaches the selection, and scores the classification terms over the highlight and over the full input.

  2. Both optimizers step, since the penalties on the mask are shared and reach the generator here as well.

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

  4. 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

PhasedSPP.training_step

Generator stepped alone in the generator phase

PhasedSPP.training_step

Detached selection for the predictor phase

PhasedSPP.phase_forward

The highlight pass

MCD.phase_class_logits

The full-input pass

MCD.extra_logits, which exposes full_class_logits

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

GRU_MCD

GRUMCDConfig

TRANSFORMER_MCD

TransformerMCDConfig

Three loss lists rather than one, and each names a phase.

Field

Default

shared_losses

sparsity and contiguity, scored in both phases

predictor_losses

classification over the highlight, and classification over the full input

generator_losses

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: 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: <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 exposes full_class_logits beside 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]])

extra_logits(input_data, selection)[source]#

The full input, which is what the highlight has to agree with.

Return type:

Dict[str, Tensor]

Parameters:
  • input_data (InputData)

  • selection (Tensor)

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

The highlight pass, which MCD trains on: in the graph, per phase.

Return type:

Tensor

Parameters:
  • input_data (InputData)

  • highlight_mask (Tensor)

  • selection (Tensor)

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

Bases: SPP

A 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_losses are scored in both phases, predictor_losses in the predictor phase and generator_losses in 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_causal2 of <jugechengzi/Rationalization-MCD> adds the sparsity and continuity terms to its classification loss and steps opt_gen beside opt_pred, and train_util.train_adv_causal of <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 registers shared_losses empty, which is what turning that flag off amounts to.

And each phase draws its own selection: phase_forward selects once per phase, so the two phases of a batch optimize different masks of it. That is the references again, which call get_rationale in each phase over a gumbel_softmax with hard=True.

The training loss is therefore not the validation loss. 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_losses empty, the generator receives none in the predictor phase. Such a model needs strategy="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_mask is what the generator produced and selection is 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.