Skip to content

[Bug] Padding positions count as negatives in the replaced-token detection loss #25

Description

@gss10282025

Summary

The RTD loss includes padding positions as negative examples. Padding the same batch to a larger width therefore changes the loss and the gradient, even though the real tokens and the replacement decisions are unchanged.

In diffcse/models.py, replaced = (g_pred != input_ids) * attention_mask gives padding positions label 0, which is a valid class, and the loss is then taken over every position:

loss_fct = nn.CrossEntropyLoss()
...
masked_lm_loss = loss_fct(prediction_scores.view(-1, 2), e_labels.view(-1))

Only the accuracy metric multiplies by attention_mask; the loss does not. Padding positions enter both the loss sum and the mean's denominator.

Environment

Commit 33b29a38, PyTorch 2.8.0+cu128, RTX 5090. The numbers below are in FP64.

Steps to reproduce

Run the official train.py and CLTrainer for one AdamW update with a randomly initialized one-layer BERT (hidden size 32), four two-view examples, dropout off, RTD weight 0.005, lr 0.001, no clipping. Pad to width 8 and to width 16 with the same real token ids, the same replacement decisions and the same initial state. Keep the custom collator in both runs so the RTD loss is actually computed.

Width 8 vs. width 16 Current code Padding labels ignored
Total loss, normalized L2 1.99930e-4 4.79383e-16
Full gradient, normalized L2 1.04826645 8.90219e-13
Parameter update, normalized L2 0.79828865 4.46219e-12

A batch that needs no padding gives the same result before and after the fix. Reproduced on a second host.

Expected behavior

The RTD loss should depend only on real tokens.

Proposed fix

Give padding positions the cross-entropy ignore index:

masked_lm_loss = loss_fct(
    prediction_scores.view(-1, 2),
    e_labels.masked_fill(attention_mask == 0, -100).view(-1),
)

This keeps the mean reduction but averages over valid tokens only. The tight numbers above are from an FP64 run; my earlier FP32 runs did not meet the same tolerances and are not presented as passing. I did not evaluate STS or try to reproduce the paper's tables.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions