G-RAT: guidance-based rationalization#
Hu and Yu, 2024, Learning Robust Rationales for Model Explainability: A Guidance-Based Approach, AAAI 2024, pages 18243-18251. Reference implementation: shuaibo919/g-rat.
Problem#
The architectures before this one fight interlocking with the signal the pair already has, which is the label. G-RAT observes that the label is a weak teacher for a selection: it says whether the words kept were sufficient, and never which words were worth keeping in the first place. A model reading the whole document can answer the second question, and the paper’s move is to build one and let it teach.
Method#
A soft attention classifier over the full input is pretrained, then keeps training beside the rationalizer, and it guides the selection in two ways at once. While its attention over the document supervises the selection directly, giving the selector a per-word target rather than a single label, its class distribution is matched against the rationalizer’s so the two models agree about the document as well as about the words. Since the guider reads everything, it is free of the bottleneck the rationalizer trains under, and its attention is a signal no highlight-only model can produce for itself.
%%{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"] --> GD["guider<br/>attention over every word"]
GD -- "attention target" --> SEL["selector"]
GD -- "class distribution" --> J{{"JSD term"}}
X --> SEL --> H["highlight h"]
H --> P["predictor"] --> Y["label"]
P --> J
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 guide 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 GD guide
class J term
The two guidance terms are annealed against each other rather than both applied throughout.
Here \(t\) counts the steps the rationalizer has taken. The first two steps both train at \(\alpha_t = 1\), as in the reference implementation, whose annealer applies its decay before counting the step. Early in training the attention target carries the weight, since the selector has nothing of its own to go on; as \(\alpha_t\) decays the distribution-matching term takes over, so the guider stops dictating the words and keeps agreeing about the label.
Training#
Two optimizers, and a rationalizer that sits out the first epochs.
The guider is stepped first, on its own classification loss over the full input.
The rationalizer selects and predicts, and the guider is read again in evaluation mode and under
no_grad, so the attention it supplies is a target rather than a path for gradient.The guider’s attention is folded onto the selection axis and normalised into a per-word target.
The five rationalizer terms are scored, with the guide and JSD terms scaled by the annealing factor.
The rationalizer’s optimizer steps only once
current_epoch >= pretrain_epochs, so the first epochs train the guider alone.
That last point is the reason warmup_epochs exists on the model.
Until the rationalizer has taken a step the monitored quantities describe a model that has not moved, and a monitor counting from epoch zero can stop a run inside that window.
Implementation#
What |
Where |
|---|---|
The guider contract |
|
The attention classifier itself |
|
The annealing factor |
|
Attention folded onto the selection axis |
|
The per-word target |
|
The staged loop |
|
Epochs the monitors should skip |
|
Gradients averaged across processes |
|
The guider attends over subtokens, since that is what its encoder reads, while the selection it guides is over words, so each word takes the attention its subtokens hold between them.
The fold sums rather than averages, because attention is a distribution and averaging would report a long word as less attended than the short one beside it.
AttentionGuider adds positive noise to its scores during training, which is the reference implementation’s way of keeping the attention from collapsing onto a handful of words.
Each of the two backward passes reaches one half of the model: the guider’s pass reaches the guider, and the rationalizer’s pass reaches the rationalizer.
A data-parallel wrapper expects every parameter it registered in every pass, so backward_and_average blocks its synchronisation for the pass.
The method then averages the gradients of the half that pass trains across processes itself.
Differences from the reference implementation#
Two details differ from the reference, and both remove a dependence on padding. Numbers from this implementation therefore do not reproduce the reference’s exactly.
The per-word target divides each word’s attention by the mean attention plus \(1 / (1 + n)\), over the \(n\) valid words. This implementation takes that mean over the valid words. The reference takes it over the padded width, so its target depends on the widest document in the batch. For a document of 10 words in a batch padded to 50, the reference divides by 0.111 and this implementation by 0.191.
The guide term is a binary cross entropy over valid words. The reference averages it over every position, padding included, where the target is zero.
Configuration#
Key |
Configuration |
|---|---|
|
|
|
|
|
|
|
|
Five loss terms rather than three, and a second list for the guider.
from pyhighlights.configurations.keys import GRU_GRAT, TOY_TASK
Registry.from_key(TOY_TASK, model=GRU_GRAT, save_path="results", seeds=[0, 1])
Registry.from_key(GRU_GRAT, pretrain_epochs=20, guide_decay=1e-3)
pretrain_epochs defaults to 10 and guide_decay to 1e-4, which is a slow anneal: the guide term still carries most of the weight after a thousand steps.
guide_loss and jsd_loss name the two terms the annealing scales, so a study registering its own criteria under different names says which ones they are.
API#
- class pyhighlights.components.models.spp.grat.AttentionGuider(backbone, predictor, noise_sigma=1.0)[source]#
Bases:
GRATGuiderBackbone-independent attention classifier used to guide G-RAT.
- Parameters:
backbone (RegistrationKey[SPPBackbone])
predictor (RegistrationKey[SPPPredictor])
noise_sigma (float)
- class pyhighlights.components.models.spp.grat.GRAT(selector_backbones, selectors, predictor, predictor_backbone, guider, guider_losses, pretrain_epochs=10, guide_decay=0.0001, guide_loss='guide', jsd_loss='jsd', **kwargs)[source]#
Bases:
SPPGuider-regularized rationalizer with staged optimization.
A soft attention classifier over the full input is pretrained, then keeps training alongside the rationalizer: its attention supervises the selection and its predictions are matched in distribution, so the generator is guided instead of regularized ad hoc.
lossesscores the rationalizer over a namespace holding the guider fields (selection_logits,guide_target,guider_class_logits) next to the model ones;guider_lossesscores the guider alone. The guide and JSD terms are annealed against each other by name.Hu and Yu, 2024, Learning Robust Rationales for Model Explainability: A Guidance-Based Approach, AAAI 2024, 18243-18251. Reference implementation: <shuaibo919/g-rat>.
- Parameters:
selector_backbones (RegistrationKey[SPPBackbone])
selectors (RegistrationKey[SPPSelector])
predictor (RegistrationKey[SPPPredictor])
predictor_backbone (RegistrationKey[SPPBackbone] | None)
guider (RegistrationKey[GRATGuider])
guider_losses (List[RegistrationKey[Loss]])
pretrain_epochs (int)
guide_decay (float)
guide_loss (str)
jsd_loss (str)
- backward_and_average(loss, parameters)[source]#
Backpropagate
lossand average the gradients ofparameters.Each of G-RAT’s two backward passes reaches one half of the model, so a data-parallel wrapper, which expects every registered parameter in every pass, would raise. The wrapper’s synchronisation is blocked for the pass, and the half that pass trains is averaged across processes here instead. On one process the average is the identity.
- Return type:
None- Parameters:
loss (Tensor)
parameters (Iterable[Parameter])
- property guide_factor: Tensor#
The guide term’s weight,
max(1 - max(t - 1, 0) * guide_decay, 0).tcounts the rationalizer’s steps. The reference implementation’sFactorAnnealerapplies its decay before counting the step, so the first two steps both train at a weight of one. A tensor rather than a float, so reading it does not synchronise with the device.
- guide_target(attention, mask)[source]#
The per-word target the guide term trains the selection towards.
min(a_i / (mean(a) + 1 / (1 + n)), 1)over thenvalid words, zero elsewhere. The mean is taken over valid words. The reference implementation takes it over the padded width, which makes a target depend on the widest document in its batch.- Return type:
Tensor- Parameters:
attention (Tensor)
mask (Tensor)
- guider_encoder_ids()[source]#
The guider’s own encoder, which is pretrained when the model’s is.
encoder_idscovers the backbones a rationalizer reads with. The guider holds a third, and a rate meant for pretrained encoders that skipped it would fine-tune one of the three at the selector’s rate.- Return type:
Set[int]
- to_selection_axis(attention, data)[source]#
The guider’s attention on the axis the selection is scored on.
The guider attends over subtokens because that is what its encoder reads. The selection it guides is over words, so each word takes the attention its subtokens hold between them, summed, since attention is a distribution. A model selecting over subtokens, or a vocabulary tokenizer, needs no folding and gets none.
- Return type:
Tensor- Parameters:
attention (Tensor)
data (InputData)
- property warmup_epochs: int#
The guider’s pretraining, which the rationalizer sits out.
training_stepgates the model’s optimizer oncurrent_epoch >= pretrain_epochs, so until then the monitored quantities describe a model that has taken no step, and a monitor counting from epoch zero can stop a run inside this window.
- class pyhighlights.components.models.spp.grat.GRATGuider(*args, **kwargs)[source]#
Bases:
Module,ABCAn attention classifier over the full input, which guides the selector.
- Parameters:
args (Any)
kwargs (Any)
- abstractmethod forward(data, mask)[source]#
Return normalized token attention [B, T] and logits [B, C].
maskis what the guider’s encoder attends over, on the encoder axis. G-RAT passesencoder_mask, which is the only axis a subword backbone can encode, and folds the attention back onto the selection axis itself.- Return type:
- Parameters:
data (InputData)
mask (Tensor)
- class pyhighlights.components.models.spp.grat.GRATGuiderOutput(attention, class_logits)[source]#
Bases:
objectWhat a guider reads off the full input.
- Parameters:
attention (Tensor)
class_logits (Tensor)
- attention: Tensor#
attention over the encoder axis, summing to one per row.
- Type:
[B, T]
- class_logits: Tensor#
the guider’s class logits.
- Type:
[B, C]
G-RAT: an attention guider regularizing the selector it is trained beside.