Accept x of shape (B, T, D). Project once into Q, K, V of (B, T, H*D_head); then reshape to (B, H, T, D_head) to keep heads independent.
scores = (Q @ K.transpose(-2, -1)) / sqrt(D_head) # (B, H, T, T)
scores = scores.masked_fill(mask == 0, float("-inf")) # mask BEFORE softmax
attn = softmax(scores, dim=-1)
out = (attn @ V) # reshape back to (B, T, D)
Gotchas separating strong from weak: (1) failing to scale by 1/sqrt(D_head) causes softmax saturation and zero-grad on first/last tokens; (2) applying softmax before the mask silently masks the wrong positions; (3) collapsing H into the batch dim instead of reshaping loses parallelism semantics; (4) PyTorch's softmax is numerically stable only along the right axis. Mention causal mask efficiency with torch.triu(torch.ones(T, T), diagonal=1) if time allows.