Skip to content

ELECTRA replaced-token detection

ModuleL6.3 · build · Python · Pass 5 · 3 to 4 h, plus your graded tests (rung R5)
You buildpython/tinyllm/obj/electra.py: ElectraDiscriminatorHead, ElectraDiscriminator, ELECTRA, replace_tokens, rtd_loss, electra_step, rtd_accuracy, save_electra, load_electra; and your own oracle tests in python/tests/l6-3-electra/
Contractcourse/contracts/py/tinyllm/obj/electra.pyi
Testscourse/tests/L6.3/test_electra.py (what they check: section 4); the oracle is Hugging Face’s ElectraForPreTraining; your tests are graded by mutation, threshold 0.80 with every required pitfall fault killed
NeedsL6.2 BertEncoder, BertForMLM, mlm_mask · M07.1 sample_categorical · L0.3 bce_with_logits · L0.1 · L0.2 · L0.4 · L0.6 safetensors · M06.3 PCG32 · M07.3 normal_init (or --ref-deps)
Used byL6.7 the zoo’s ELECTRA row (replaced-token detection accuracy) · the discriminator body is the alternative classification backbone of L6.5 (tl_arch = "electra")
MilestoneMS-L6 (train electra, then a LoRA classifier over its discriminator)
Optional depthClark, Luong, Le, and Manning, “ELECTRA: Pre-training Text Encoders as Discriminators Rather Than Generators” (ICLR 2020), sections 2, 3.2, and appendix A; Goodfellow et al., “Generative Adversarial Nets” (2014), for what ELECTRA is not
  • The discriminator learns from every real token, not only the 15% BERT masks: each one is labelled original or replaced (test_electra_step_is_its_pieces).
  • A sampled token equal to the original is labelled original (test_a_lucky_sample_is_original, test_hand_example_replaced_token_labels).
  • The sample is data: the generator trains on its MLM loss only, and the total is LG+λLDL_G + \lambda L_D with λ=50\lambda = 50 (test_no_gradient_through_the_sample, test_lambda_weights_the_discriminator).
  • The discriminator is BERT’s encoder plus a dense, GELU, dense head, and its loss averages over real tokens only; both match Hugging Face (test_discriminator_matches_hf, test_rtd_loss_ignores_padding).
  • Generator and discriminator share one set of embeddings (test_tied_embeddings_are_one_tensor).
Terminal window
ol start L6.3 # stubs electra.py; prints your test path and rung (R5)
ol tests L6.3 # the course tests
# write your oracle tests in python/tests/l6-3-electra/, then:
ol check L6.3 # course tests and the mutation grade of your tests
ol mutate L6.3 # the full grade, cached by your test files' hash
ol diff L6.3 # after passing: your code against the reference

Your BERT (L6.2) learns from a masked-LM loss on 15% of the tokens of each batch; the other 85% cost a full forward and backward pass and teach nothing. On the small corpus and compute budget of this course that waste shows: a BERT pretrained for the MS-L6 step budget is a weak backbone for the sentiment classifier of L6.5. ELECTRA turns every position into a training signal with the same encoder: a small generator fills the masked positions with plausible guesses, and the encoder, now a discriminator, says for every token whether it was replaced. The zoo (L6.7) puts the two backbones side by side on the same classification task, which is the comparison the paper is about.

SymbolMeaningType / shape
xxthe batch of token idsint64[B, T]
mmthe real-token mask (attn_mask), 1 = real, 0 = paddingbool[B, T]
x~\tilde x, ℓ\ellthe masked inputs and MLM labels from mlm_mask (-100 = not picked)int64[B, T]
GGthe generator, a small BertForMLMModule
DDthe discriminator, BertEncoder plus a token headModule
pG(⋅∣x~,t)p_G(\cdot \mid \tilde x, t)the generator’s distribution at position ttfloat64[V]
x^\hat xthe corrupted ids (corrupt)int64[B, T]
yt=[x^t≠xt]y_t = [\hat x_t \ne x_t]the discriminator’s label (is_replaced)float32[B, T]
LGL_G, LDL_DMLM cross-entropy, replaced-token BCEscalars
λ\lambdathe discriminator’s loss weight, 50float
  1. Mask like BERT: $(\tilde x, \ell) = $ mlm_mask(x, special ∨¬m\lor \lnot m , mask_id, V, p, rng). Padding counts as special, so it is never picked.
  2. Run the generator on x~\tilde x: logits and LGL_G, the mean cross-entropy over the picked positions.
  3. At each picked position, in C order, draw one uniform and sample x^t∼pG(⋅∣x~,t)\hat x_t \sim p_G(\cdot \mid \tilde x, t) with sample_categorical (M07.1, inverse CDF in float64); elsewhere x^t=xt\hat x_t = x_t.
  4. Label every token: yt=1y_t = 1 when x^t≠xt\hat x_t \ne x_t. Run the discriminator on x^\hat x and take LDL_D, the mean binary cross-entropy over the real tokens.
  5. L=LG+λLDL = L_G + \lambda L_D.

The random draws are mlm_mask’s, then one uniform per picked position, all from one generator. Sampling is not differentiable, and the paper does not try to make it so (an adversarial generator trained by reinforcement learning did worse): the corrupted ids enter the discriminator as plain data, so LDL_D sends no gradient into GG. Each network learns from its own loss, and one backward pass of LL does both.

The label asks “is this token different from the original?”, not “was this position masked?”. A good generator often samples the original token back, especially for function words; those positions are labelled original. Labelling every masked position as replaced teaches the discriminator to flag tokens that never changed, and its accuracy then measures the masking rate, not the text.

The body is L6.2’s BertEncoder, unchanged, under the name electra. The head is Hugging Face’s ElectraDiscriminatorPredictions: dense (d→dd \to d), the exact GELU, dense_prediction (d→1d \to 1), and the last axis dropped, so the output is one logit per token, positive meaning “replaced”. The loss is L0.3’s bce_with_logits over the real tokens only, the mean over them: padding contributes neither to the sum nor to the count. With these names, a Hugging Face ELECTRA checkpoint loads through L6.2’s hf_bert_encoder_sd after dropping its electra. prefix, and the golden test does exactly that.

ELECTRA ties the generator’s token, position, and type embeddings to the discriminator’s: one Tensor under two names. The generator is otherwise smaller (fewer layers here; the paper also narrows it), which matters: a generator as strong as the discriminator produces replacements too hard to detect. In the ELECTRA module the discriminator registers first, so the shared Tensors are listed once, under discriminator., and both losses add their gradients into them.

The model directory holds both networks, the tie, the mask id, and the special ids, so the zoo can corrupt held-out text the same way and report replaced-token detection accuracy (rtd_accuracy): over the real tokens, how often the discriminator’s call (logit >0> 0) matches yy.

Four tokens x=(5,6,7,8)x = (5, 6, 7, 8); mlm_mask picked positions 1 and 3, so ℓ=(−100,6,−100,8)\ell = (-100, 6, -100, 8). The generator’s distributions at the picked positions:

positionpGp_Guniformcumulative sum crosses it atx^t\hat x_tyty_t
10.75 on token 6, 0.25 on token 90.1token 6 (0.75)60 (the original)
30.5 on token 2, 0.5 on token 80.3token 2 (0.5)21

So x^=(5,6,7,2)\hat x = (5, 6, 7, 2) and y=(0,0,0,1)y = (0, 0, 0, 1). Position 1 was masked, but the sample is the original token: not replaced.

With all discriminator logits 0, each token’s BCE is ln⁡2\ln 2 and LD=ln⁡2=0.693147L_D = \ln 2 = 0.693147. With logits (−2,−2,−2,2)(-2, -2, -2, 2) every call is right and each token pays ln⁡(1+e−2)=0.126928\ln(1 + e^{-2}) = 0.126928, the mean too.

This is test_hand_example_replaced_token_labels.

class ElectraDiscriminator(Module): # electra.* (BertEncoder) + discriminator_predictions.*
def __init__(self, cfg: BertConfig, rng=None) -> None: ...
def forward(self, ids, token_type_ids=None, attn_mask=None) -> Tensor: ... # [B, T] logits
class ELECTRA(Module): # discriminator.*, generator.* (tied embeddings)
def __init__(self, gen_cfg: BertConfig, disc_cfg: BertConfig, rng=None, tie_embeddings=True) -> None: ...
def replace_tokens(ids, labels, gen_logits, rng, ignore_index=-100) -> tuple[NDArray, NDArray]: ...
def rtd_loss(disc_logits: Tensor, is_replaced, attn_mask=None) -> Tensor: ...
def electra_step(gen, disc, ids, rng, lam=50.0, *, mask_id, vocab=None, special_mask=None,
token_type_ids=None, attn_mask=None, p=0.15) -> dict: ...
def rtd_accuracy(model: ELECTRA, ids, rng, *, mask_id, ...) -> NDArray: ...
def save_electra(model, dir, mask_id, special_ids=(), tokenizer="file") -> None: ...
def load_electra(dir) -> tuple[ELECTRA, dict]: ...
TestKINDChecksWhy it matters downstream
test_hand_example_replaced_token_labelsunitsection 3: x^\hat x, yy, two uniforms used, both loss valuesyou and the test agree on the labels
test_a_lucky_sample_is_originalboundarya certain generator replaces nothingthe label is “different”, not “masked”
test_replace_draws_one_uniform_per_masked_positionunitthree picked positions, three uniforms, C orderthe same seed gives the same corruption everywhere
test_discriminator_matches_hfgoldenlogits, loss, and gradients of HF’s ElectraForPreTrainingpretrained ELECTRA checkpoints load and score alike
test_rtd_loss_ignores_paddingpropertypadded logits and labels change nothingpadded batches train the same as unpadded ones
test_electra_step_is_its_piecesdifferentialthe step equals mask, generate, sample, discriminate, recomputed from a copied rngthe order of the five steps
test_padding_is_never_masked_or_scoredpropertypadding and specials are never picked or replacedno learning from padding
test_no_gradient_through_the_samplepropertygenerator gradients of LL equal those of LGL_G aloneeach network learns from its own loss
test_lambda_weights_the_discriminatorunitdefault λ=50\lambda = 50; L=LG+λLDL = L_G + \lambda L_D for three valuesthe paper’s balance of the two losses
test_tied_embeddings_are_one_tensorpropertyone Tensor, listed once under discriminator.; mismatched configs raiseshared embeddings, saved once
test_save_load_roundtrippropertythe directory restores both networks, the tie, the mask id, and the specialsthe zoo loads it
test_rtd_accuracy_definitionunitan “always original” discriminator scores the fraction not replacedthe zoo’s ELECTRA metric

Your oracles: hand-made generator logits whose samples you know (a certain token, a uniform row and chosen uniforms), the BCE written out with math.log1p, and the step recomputed from its pieces with a deep copy of the generator. Cover the lucky sample, one draw per picked position, the head’s formula (dense, exact GELU, dense) recomputed in numpy, the real-token mean, padding never picked, the discriminator reading the corrupted ids, and the generator’s gradients coming from LGL_G only. Import only the contract (tinyllm.obj.electra, tinyllm.obj.bert). ol check L6.3 requires a mutation score of at least 0.80 with every required pitfall fault killed.

PitfallSymptomCaught by
1. every masked position labelled replacedthe discriminator learns where the masks were; its accuracy tracks the masking ratetest_a_lucky_sample_is_original, test_hand_example_replaced_token_labels (mutant s01)
2. a uniform drawn at every positiona seed gives another corruption than the spec, so Python and a replay disagreetest_replace_draws_one_uniform_per_masked_position (mutant s02)
3. a head without the GELUlogits differ from HF; pretrained checkpoints score differentlytest_discriminator_matches_hf (mutant s03)
4. padding in the loss or in the maskingshort sequences in a batch weigh less; padding gets replacedtest_rtd_loss_ignores_padding (mutant s04), test_padding_is_never_masked_or_scored (mutant s05)
5. the discriminator reading the masked inputsit sees [MASK] where the replacement should be: the task becomes “find the masks”test_electra_step_is_its_pieces (mutant s06)
6. λ\lambda on the wrong lossthe generator’s MLM loss dominates and its gradients grow 50 timestest_lambda_weights_the_discriminator, test_no_gradient_through_the_sample (mutant s07)
7. untied embeddingstwo embedding tables, the generator’s trained only by LGL_Gtest_tied_embeddings_are_one_tensor (mutant s08)
DirectionModuleHow it uses this
BackL6.2mlm_mask picks positions; BertForMLM is the generator and BertEncoder the discriminator’s body
BackM07.1sample_categorical draws the replacements
BackL0.3bce_with_logits is the replaced-token loss
BackL0.1the Tensor and no_grad (the zoo scores without a graph)
BackL0.2gelu and reshape in the head
BackL0.4the head’s Linear layers and the Module registration order
BackL0.6the model directory’s safetensors
BackM06.3PCG32: the default init stream and every draw of the step
BackM07.3normal_init for the head’s weights
ForwardL6.7the zoo loads ELECTRA directories and reports replaced-token detection accuracy next to BERT’s masked-token accuracy

If you skip this module, ol check L6.7 stops with needs L6.3: build it, or pass --ref-deps.

Your pieceProduction equivalentWhat it addsWhere to look
ElectraDiscriminatorHF ElectraForPreTrainingembeddings_project when the embedding size differs from the hidden sizetransformers/models/electra/modeling_electra.py
electra_stepthe original ELECTRA pretraining codea generator a third to a quarter of the discriminator’s width; gumbel-noise sampling; masking 85% [MASK], 15% keptgoogle-research/electra, pretrain/pretrain_helpers.py
replaced-token detectionDeBERTaV3gradient-disentangled embedding sharing: the generator’s gradients stop at the shared embeddingsHe, Gao, and Chen, “DeBERTaV3” (2021)
ELECTRA as a backboneELECTRA-small fine-tuned on GLUESST-2 and the other GLUE tasks with a classification headL6.5; the paper’s table 1