Train a small event-sequence decoder with causal and padding controls, then break each control deliberately to prove the tests detect leakage.
Project: audit a next-event attention decoder
Construct event sequences
Generate 470 customer-service event histories with a fixed vocabulary, a start event, an end event and right padding. Split by customer identity before forming overlapping windows so nearby windows from one customer cannot land in both train and validation. Record the tokenizer version, maximum sequence length and count of valid target positions. Grouped temporal validation keeps the evaluation population honest.
Implement one attention block
Project input embeddings into queries, keys and values, preserve batch and sequence axes, and combine causal and valid-key masks. Exclude padded query positions from pooling or downstream metrics. Pass zero attention dropout during evaluation. On a test pair that differs only after position three, compare earlier outputs exactly within a numerical tolerance. The mask contract defines the expected shape and truth values.
Align labels and loss
Train with inputs excluding the final token and targets excluding the first. Ignore padding targets; include the end-event target if the product needs a stop prediction. Sum unreduced token losses and divide by valid targets across each effective batch. Add a test where every target is ignored and reject that batch instead of dividing by zero. Target shifting is independent of the attention mask.
Run negative controls
First remove the causal mask and show that changing a future event alters an earlier output. Then keep the mask but stop shifting targets; show that a leakage assertion catches the self-target path. Finally count padding positions in the loss and report the denominator error. A model with a lower leaked loss fails the project; it is learning information unavailable when the product makes its prediction.
Publish a checked result
Provide fixed data manifests, model revision, mask tensor shapes, causal counterfactual results, valid-target counts, training and validation loss, and inference latency at the maximum accepted length. Restart from a checkpoint and compare the next update. Report how much compute was spent on padding and state a length limit. The final serving artifact must include vocabulary IDs and start/end semantics, not only weights.
Implementation
import torch
from torch import nn
import torch.nn.functional as functional
class ServiceEventAttention(nn.Module):
def __init__(self, vocabulary_size=47, width=12):
super().__init__()
self.embedding = nn.Embedding(vocabulary_size, width)
self.project_qkv = nn.Linear(width, width * 3, bias=False)
def forward(self, event_ids, valid_tokens):
if not bool(valid_tokens.any(dim=1).all()):
raise ValueError("each history needs a valid token")
hidden = self.embedding(event_ids)
query, key, value = self.project_qkv(hidden).chunk(3, dim=-1)
positions = event_ids.size(1)
causal_allow = torch.ones(positions, positions,
dtype=torch.bool, device=event_ids.device).tril()
allow_mask = (causal_allow[None, None, :, :]
& valid_tokens[:, None, None, :])
attended = functional.scaled_dot_product_attention(
query.unsqueeze(1), key.unsqueeze(1), value.unsqueeze(1),
attn_mask=allow_mask, dropout_p=0.0)
return attended.squeeze(1) * valid_tokens[:, :, None]
torch.manual_seed(83)
decoder_block = ServiceEventAttention().eval()
first_history = torch.tensor([[7, 12, 19, 23]])
changed_future = torch.tensor([[7, 12, 19, 31]])
valid_tokens = torch.ones_like(first_history, dtype=torch.bool)
with torch.inference_mode():
first_output = decoder_block(first_history, valid_tokens)
second_output = decoder_block(changed_future, valid_tokens)
torch.testing.assert_close(first_output[:, :3], second_output[:, :3])Performance and operating cost
The one-head reference uses O(B × L squared × D) attention arithmetic, plus embeddings and QKV projections. The causal mask protects information flow but does not by itself make long sequences cheap. Profile peak memory and latency at the accepted maximum length, and compare valid-token throughput with allocated-token throughput. Identity-safe validation and negative controls cost extra runs but prevent a misleading low loss.
Common Mistakes
- Do not split overlapping windows from one customer across train and validation.
- Do not take a low loss as proof that future events were hidden.
- Do not publish an inference block without the same vocabulary, length limit and stop-token policy used in training.
Read next
- Scaled dot-product attention and mask contracts
- Causal target shifts and padding-aware loss
- Checkpoint recovery: save optimizer state and the run boundary
- Text validation: split conversations, duplicates and time together
- Inference contracts: preserve preprocessing and measure tail latency
Continue the workflow: Project: predict service escalation from event histories.
Continue the workflow: Project: release an incremental service-event decoder.
