How do you translate "আমি ভাত খাই" into "I eat rice"? The input and output are both sequences, of different lengths, with different word orders. In 2014 two groups — Sutskever, Vinyals and Le at Google, and Cho et al. — showed that neural networks could learn such mappings end to end with the sequence-to-sequence (seq2seq) encoder–decoder architecture. It transformed machine translation and became a template for summarisation, speech recognition, question answering and code generation.
The architecture#
Encoder: an RNN reads the source sequence $x_1, \dots, x_S$ and produces hidden states; its final state $\mathbf{c} = \mathbf{h}_S$ is a fixed-size context vector summarising the whole input.
Decoder: another RNN, initialised with $\mathbf{c}$, generates the target one token at a time. At step $t$ it takes the previous output token $y_{t-1}$ and its state, and predicts a distribution over the next token:
Generation begins with a special start token <sos> and stops when the decoder emits <eos>.
Training with teacher forcing#
The loss is the sum of cross-entropies over target positions. During training we feed the decoder the ground-truth previous token rather than its own prediction — teacher forcing. This makes training stable and parallelisable across time steps (for the decoder inputs), but creates exposure bias: at inference the decoder sees its own, possibly wrong, predictions, a situation it never encountered during training. Scheduled sampling gradually replaces ground-truth inputs with model predictions to reduce this mismatch.
Decoding strategies#
Finding the most probable output sequence exactly is intractable (the space is exponential). Approximations:
- Greedy decoding — pick the most probable token each step. Fast, but an early mistake cannot be undone.
- Beam search — keep the $k$ best partial sequences (the beam) at each step, ranked by total log-probability; expand each by every token and keep the top $k$ again. Beam widths of 4–10 are typical in translation.
Beam search favours short sequences (each extra token adds a negative log-probability), so scores are length-normalised, e.g. divided by $T^\alpha$ with $\alpha \approx 0.6$–$1$. For open-ended generation, sampling methods (top-k, nucleus) are preferred — covered in the LLM lectures.
A compact implementation#
import torch
import torch.nn as nn
class Encoder(nn.Module):
def __init__(self, V, E=64, H=128):
super().__init__()
self.emb, self.rnn = nn.Embedding(V, E), nn.GRU(E, H, batch_first=True)
def forward(self, src):
outputs, h = self.rnn(self.emb(src))
return outputs, h # outputs kept for attention later
class Decoder(nn.Module):
def __init__(self, V, E=64, H=128):
super().__init__()
self.emb, self.rnn, self.out = nn.Embedding(V, E), nn.GRU(E, H, batch_first=True), nn.Linear(H, V)
def forward(self, tgt_in, h):
o, h = self.rnn(self.emb(tgt_in), h)
return self.out(o), h
class Seq2Seq(nn.Module):
def __init__(self, V_src, V_tgt):
super().__init__()
self.enc, self.dec = Encoder(V_src), Decoder(V_tgt)
def forward(self, src, tgt_in): # teacher forcing
_, h = self.enc(src)
logits, _ = self.dec(tgt_in, h)
return logits
@torch.no_grad()
def greedy(self, src, sos, eos, max_len=30):
_, h = self.enc(src)
y = torch.full((src.size(0), 1), sos)
out = []
for _ in range(max_len):
logits, h = self.dec(y, h)
y = logits[:, -1].argmax(-1, keepdim=True)
out.append(y)
if (y == eos).all():
break
return torch.cat(out, 1)
# Toy task: reverse a sequence of digits (tokens 3..12); 0=pad, 1=sos, 2=eos
V = 13
model = Seq2Seq(V, V); opt = torch.optim.Adam(model.parameters(), 2e-3)
for step in range(3000):
src = torch.randint(3, V, (64, 8))
tgt = torch.cat([src.flip(1), torch.full((64, 1), 2)], 1)
tgt_in = torch.cat([torch.full((64, 1), 1), tgt[:, :-1]], 1)
loss = nn.functional.cross_entropy(model(src, tgt_in).reshape(-1, V), tgt.reshape(-1))
opt.zero_grad(); loss.backward(); opt.step()
test = torch.randint(3, V, (3, 8))
print(test.tolist()); print(model.greedy(test, 1, 2).tolist())The bottleneck problem#
The entire source sentence must be compressed into a single fixed-size vector $\mathbf{c}$. For short sentences this works; for long ones, information is lost. Cho et al. (2014) observed that translation quality degraded sharply with sentence length. Sutskever et al. found a curious trick helped: reversing the source sentence, which placed the first source words close to the first target words and shortened the dependencies the network had to bridge.
The principled fix was to let the decoder look back at all encoder states, choosing which to focus on at each step. That is attention — the subject of the next lecture, and the seed from which the Transformer grew.