A next-token learner should predict position t plus one from positions no later than t, while padding and any unavailable targets contribute no loss.
Causal target shifts and padding-aware loss
Separate model input from target
For tokens A, B, C and D, feed A, B and C and score predictions against B, C and D. If the model receives D at the position used to predict D, the task leaks its own answer even when the attention mask is causal. Make the one-position shift explicit in the data contract, and assert input length equals target length after shifting. Tokenization fixes the vocabulary and special-token policy before this calculation.
Mask the target, not just attention
A key padding mask can keep a token from being attended to but does not stop a padded target from entering cross-entropy. Replace target positions without a real next token by the loss function’s ignore value, or use a weighted unreduced loss with a verified denominator. A batch with one long sequence and many short sequences can otherwise be dominated by padding. Count valid target tokens per optimizer update and report loss per valid token, not per allocated slot.
Test causal independence
Hold a prefix fixed, change only a future token, and assert that logits for earlier positions stay unchanged under deterministic evaluation. This catches a reversed triangular mask, an accidental bidirectional layer and shifted labels that read future information. A model can still score well on a leaky validation set, so the counterfactual test is more decisive than a low loss. Attention mask semantics determine which comparisons are legal.
Handle boundary tokens deliberately
A start token can give the first target a predecessor; an end token lets the model learn when to stop. Whether the final input predicts an end token is a modeling choice, not a side effect of padding. Left-padding during generation differs from right-padding during training; position IDs and cached key/value length may need adjustment. Document the serving prompt limit and stop behavior alongside the training shift. Inference contracts protect this boundary.
Measure the actual denominator
For each batch, log allocated positions, valid input positions and valid target positions. Reject an all-padding batch before softmax or loss reduction. If using gradient accumulation with variable lengths, sum token losses and divide by the total valid targets across the effective batch; averaging per-sequence means first gives short sequences extra weight. Pin the token weighting rule before comparing runs.
Implementation
import torch
import torch.nn.functional as functional
pad_id = 0
event_sequences = torch.tensor([[7, 12, 19, 23, 0],
[7, 31, 17, 0, 0]])
model_inputs = event_sequences[:, :-1]
next_targets = event_sequences[:, 1:].clone()
next_targets[next_targets == pad_id] = -100
torch.manual_seed(29)
vocabulary_size = 47
next_token_logits = torch.randn(2, 4, vocabulary_size,
requires_grad=True)
token_loss_sum = functional.cross_entropy(
next_token_logits.transpose(1, 2), next_targets,
ignore_index=-100, reduction="sum")
valid_targets = (next_targets != -100).sum()
mean_token_loss = token_loss_sum / valid_targets
mean_token_loss.backward()
assert int(valid_targets) == 5
assert model_inputs.shape == next_targets.shape
assert torch.isfinite(mean_token_loss)Performance and operating cost
Shifting itself costs O(B × L) views or copies depending on later operations. A vocabulary cross-entropy over V classes costs roughly O(B × L × V) arithmetic and logits storage unless a fused implementation changes intermediates. Padding consumes attention and output compute even when ignored by loss. Report valid-token throughput, not only padded tokens per second.
Common Mistakes
- Do not score each token against itself because the labels were left unshifted.
- Do not assume an attention mask removes padded positions from the loss.
- Do not average unequal-length sequence means when the intended unit is one valid target token.
Read next
- Scaled dot-product attention and mask contracts
- Logits, cross-entropy and gradients: align the training calculation
- Unicode and tokenization: preserve meaning at the text boundary
- Gradient accumulation and effective batch accounting
- Project: audit a next-event attention decoder
Continue the workflow: Packed LSTM sequences and the last valid state.
Continue the workflow: Decoding stop rules, sampling and beam-score accounting.
Continue the workflow: CTC time lengths, blank labels and alignment feasibility.
Continue the workflow: Causal dilated convolution and receptive field.
Continue the workflow: Teacher forcing, target shifts and generation parity.
