DAR: discriminatively aligned rationalization#

Liu, Wang, Wang, Deng, Zhang, Wang and Li, 2024, Enhancing the Rationale-Input Alignment for Self-explaining Rationalization, ICDE 2024, pages 2218-2230. Reference implementation: jugechengzi/dar.

Problem#

Selector and predictor are trained together on one signal, so nothing stops them from agreeing on a private code. The selection drifts away from the semantics of the document, and the predictor learns to read the drift. Accuracy stays high, and the generator is rewarded for a highlight nobody else can interpret. The paper calls that rationale shift, and it is the failure mode where every reported number looks correct and the highlight is unreadable.

Method#

DAR adds a third module that never learned the code. An aligner, which is a second predictor, is trained on the full input alone and then frozen. Its only notion of what a label looks like therefore comes from ordinary text. Since that module has never seen a highlight during its own training, it can only read one the way it reads text. Asking it to predict the label from the highlight therefore costs the generator anything it selected in a private code.

        %%{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 --> P["predictor<br/>trained with the generator"] --> Y["label"]
    X -- "pretraining, then frozen" --> A["aligner<br/>reads full text only"]
    H --> A
    A --> AL{{"alignment term<br/>trains the generator"}}

    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 frozen fill:#e8e4f3,stroke:#4a3b76,stroke-width:2px,color:#241a3d
    classDef term fill:#ffe9c7,stroke:#a35c00,stroke-width:2px,color:#3d2100
    class SEL,P head
    class X,H,Y value
    class A frozen
    class AL term
    
\[\mathcal{L} = \underbrace{\mathcal{L}_{\text{cls}}\big(p_\phi(y \mid h \odot x),\, y\big)}_{\text{the pair, as usual}} + \lambda_s \Omega_s(h) + \lambda_c \Omega_c(h) + \underbrace{\mathcal{L}_{\text{cls}}\big(p_{\bar\psi}(y \mid h \odot x),\, y\big)}_{\text{the frozen aligner}}\]

The bar over \(\psi\) is the point of the method: the aligner’s parameters do not move while the rationalizer trains, so the fourth term trains the generator and nothing else. An aligner that kept moving could co-adapt to the highlight, which is the one thing it exists not to do.

Training#

Pretraining first, then an ordinary single-optimizer loop.

  1. Before the first rationalization epoch, the aligner is trained on the full input for pretrain_epochs epochs, in a loop the model drives itself. After each epoch, its F1 on the full input of the validation split is computed, and the aligner of the best epoch is kept, as in the reference implementation.

  2. Its parameters are frozen and it is set to evaluation mode. A flag in a buffer records this, so a resumed run finds it trained rather than pretraining a second one.

  3. Every batch afterwards selects, predicts from the highlight, and asks the frozen aligner to predict from the same highlight.

  4. The four terms are summed and one backward pass updates the generator and the predictor, never the aligner.

Implementation#

What

Where

The pretraining loop

DAR.pretrain_aligner, called from on_train_start

The aligner reading a highlight

DAR.align

The aligner reading full text

DAR.align_full

The aligner left out of the optimizer

DAR.configure_optimizers

Pretraining kept across a checkpoint

the aligner_ready buffer

The alignment term

DAR.compute_loss, which exposes aligner_class_logits

The aligner is pretrained by the model rather than by the task, so it is part of the model, is checkpointed with it, and a resumed run finds it trained. The pretraining loop runs outside the strategy Lightning drives. Under data parallelism it averages each gradient across processes itself, so every process holds the same aligner. An EarlyStopping callback counts epochs of the rationalizer rather than of the pretraining, which keeps a long pretraining off the patience counter.

Differences from the reference implementation#

The reference implementation reloads its pretrained aligner in training mode. Its dropout therefore runs on every highlight the aligner scores, and the alignment term changes between two readings of the same highlight. This library treats that as an error in the reference. Once frozen, the aligner stays in evaluation mode, and the alignment term is a fixed function of the highlight.

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_DAR

GRUDARConfig

TRANSFORMER_DAR

TransformerDARConfig

from pyhighlights.configurations.keys import GRU_DAR, TOY_TASK

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

Registry.from_key(GRU_DAR, pretrain_epochs=50)

pretrain_epochs sets the length of the pretraining. The aligner of the epoch with the highest validation F1 is kept: the F1 of class 1 for two classes, and the macro F1 for more. Without a validation loader, or with limit_val_batches=0, the aligner of the last epoch is kept. aligner_backbone is the aligner’s own encoder and its head is built from predictor, since the two modules answer the same question about the same labels. The alignment term is appended to losses by the model rather than declared in the list, so a registration cannot forget the term the method is.

API#

class pyhighlights.components.models.spp.dar.DAR(aligner_backbone, predictor, aligner_loss, pretrain_epochs=20, **kwargs)[source]#

Bases: SPP

Rationalizer whose highlight has to read like the input it came from.

A cooperative game lets the pair agree on a private code. The selection drifts away from the semantics of the full input, and the predictor learns to read the drift. Accuracy stays high, and the generator is rewarded for a highlight nobody else can interpret. The paper calls that rationale shift. DAR answers it with a second predictor, the aligner, trained on the full input alone and then frozen. That module never sees a highlight during its own training, so it can only read one the way it reads text. Asking it to predict the label from the highlight therefore costs the generator anything it selected in a private code. The aligner is frozen, so this term trains the generator only.

Liu, Wang, Wang, Deng, Zhang, Wang and Li, 2024, Enhancing the Rationale-Input Alignment for Self-explaining Rationalization, ICDE 2024, pages 2218-2230. Paper: <https://doi.org/10.1109/ICDE60146.2024.00176>. Reference implementation: <jugechengzi/dar>.

The aligner is pretrained here rather than by the task: it is part of the model, is checkpointed with it, and a resumed run finds it trained. The pretraining runs its own loop over the training loader before the first epoch, and keeps the aligner of its best validation epoch. Under data parallelism every process averages its gradients with the others’, so all processes hold the same aligner. An EarlyStopping callback counts epochs of the rationalizer only, so a long pretraining stays off the patience counter.

The frozen aligner stays in evaluation mode, so the alignment term is a fixed function of the highlight. The reference implementation reloads its aligner in training mode, so its dropout runs on every highlight the aligner scores. This library treats that as an error in the reference.

Parameters:
  • aligner_backbone (RegistrationKey[SPPBackbone])

  • predictor (RegistrationKey[SPPPredictor])

  • aligner_loss (RegistrationKey[Loss])

  • pretrain_epochs (int)

align(data, highlight_mask)[source]#

What the aligner makes of a selection.

Return type:

Tensor

Parameters:
  • data (InputData)

  • highlight_mask (Tensor)

align_full(data)[source]#

What the aligner makes of the whole input, which is all it is taught.

Return type:

Tensor

Parameters:

data (InputData)

aligner_f1(loader)[source]#

The aligner’s F1 on the full input of loader.

Two classes score the F1 of class 1, as the reference implementation does, and more classes score the macro F1. The counts are summed across processes, so every process reads the same score. A class with no predictions and no examples scores zero.

Return type:

float

pretrain_aligner()[source]#

Train the aligner on the full input, then freeze it.

Once per fit and before the first epoch, so every rationalization batch is scored against the same module. An aligner that kept moving could co-adapt to the highlight, which is the thing it exists not to do.

With a validation loader, the aligner of the epoch with the highest aligner_f1() on it is the one kept, as in the reference implementation. The first such epoch wins a tie. Without one, or with limit_val_batches=0, the aligner of the last epoch is kept.

Return type:

None

DAR: a highlight scored by a module that only ever read the full input.