Natural Language Processing Basics

Bidirectional & Stacked RNNs


You are tagging entities in news text. A recurrent model reads left to right and reaches the first word here:

Text
"Washington ..."

Is it a person or a place? You cannot know. Neither can the model. Everything that would settle it is still in the future:

Text
"Washington signed the treaty in 1796."      -> PERSON"Washington approved the new bill today."     -> could be either"Washington is humid in August."              -> LOCATION"Washington scored twice in the second half." -> PERSON (a footballer)

A left-to-right model must commit to a tag for Washington using only the words before it — of which there are none. It will guess whichever tag was more common in training and be wrong a predictable fraction of the time. No amount of extra training data fixes this, because the required information is genuinely not available at that point in the model's traversal.

The same blind spot shows up everywhere. In "the bank was steep and muddy", disambiguating bank requires steep, which arrives three words later. In "he said the food was, to be fair, inedible", the sentiment verdict sits at the end of a clause whose earlier words all read positively.

Depth stacks in one direction, width runs in twoTokens in, one embedding per positionLayer 1 forward, and layer 1 backwardConcatenate: hidden size doubles hereLayer 2 reads that whole sequencePool for a label, or tag every position
Bidirectional is off the table the moment the input arrives token by token — the backward pass needs the end first.

Reading in both directions

The fix is direct: run two independent recurrent networks over the sequence, one forwards and one backwards, and combine their states at each position.

Text
                Washington  signed    the      treatyforward   ->    hf1    ->   hf2   ->  hf3  ->  hf4backward  <-    hb1    <-   hb2   <-  hb3  <-  hb4output at position 1 = [hf1 ; hb1]                        ^      ^                        |      knows about "signed the treaty"                        knows only about "Washington"

Formally, with ⊕\oplus meaning concatenation:

h→t=RNNf(xt,h→t−1),h←t=RNNb(xt,h←t+1),ht=h→t⊕h←t\overrightarrow{h}_t = \text{RNN}_f(x_t, \overrightarrow{h}_{t-1}), \quad \overleftarrow{h}_t = \text{RNN}_b(x_t, \overleftarrow{h}_{t+1}), \quad h_t = \overrightarrow{h}_t \oplus \overleftarrow{h}_t

The two networks share nothing — separate weights, separate states. Concatenating them means the output at every position has dimension 2dh2 d_h, which is the single most common source of shape errors in bidirectional code.

Every position now has access to the entire sequence: the forward state summarises everything to the left, the backward state everything to the right. Nothing is hidden from any timestep.

The gain is not marginal. On named entity recognition, going bidirectional is typically worth 5 to 15 F1 points — one of the largest single-change improvements available in recurrent modelling.

Where you cannot use it

Bidirectionality requires the whole sequence up front. That rules out three important cases, and getting this wrong produces bugs that are invisible offline and catastrophic in production.

SituationBidirectional?Why
Classifying a complete documentYesWhole text is available
Tagging a complete sentenceYesWhole sentence is available
Encoder side of a translation modelYesSource sentence is complete
Predicting the next wordNoThe backward pass has already read the answer
Live transcription or streaming inputNoThe future does not exist yet
Any autoregressive generationNoSame leakage as next-word prediction

The leakage case is worth spelling out. If you build a language model with a bidirectional encoder, the backward state at position tt contains information about the token at position t+1t+1 — which is exactly what you are asking the model to predict. Training loss will drop to near zero and perplexity will look extraordinary. At inference, the future is unavailable, and the model produces nonsense. The symptom is a model that appears to work perfectly until the moment it has to do the job.

Stacking layers

The second structural change is depth: feed the output sequence of one recurrent layer into another as its input sequence.

Text
layer 2:   h2_1 -> h2_2 -> h2_3 -> h2_4     (reads layer 1's outputs)             ^      ^       ^       ^layer 1:   h1_1 -> h1_2 -> h1_3 -> h1_4     (reads embeddings)             ^      ^       ^       ^inputs:     x1     x2      x3      x4

Layer 1 sees word embeddings. Layer 2 sees a sequence of context-aware vectors — each already summarising a neighbourhood. This lets the model build a hierarchy, in the same way that stacked convolutional layers move from edges to textures to objects.

Probing studies on trained stacked models find a fairly consistent division of labour:

LayerTends to encode
1Morphology, part of speech, local collocations (New York, not good)
2Phrase structure, short-range syntax, entity boundaries
3+Clause-level and discourse-level meaning

Depth is not free, and the returns fall off sharply:

LayersParameters (relative)Typical accuracy changeVerdict
11.0×baselineFine for short texts and small datasets
2~2.3×+1 to +3 pointsThe usual sweet spot
3~3.6×+0 to +1 pointSometimes worth it on large datasets
4+~5×Often negativeNeeds residual connections to train at all

Beyond three layers, gradients must now travel both backwards through time and downwards through layers, and the vertical path has no gating protecting it. Residual connections — adding a layer's input to its output — give the gradient a direct route and are what make deeper stacks trainable:

Python
class ResidualLSTMStack(nn.Module):    def __init__(self, input_dim, hidden_dim, num_layers, dropout=0.3):        super().__init__()        self.layers = nn.ModuleList()        self.projections = nn.ModuleList()        dim = input_dim        for _ in range(num_layers):            self.layers.append(                nn.LSTM(dim, hidden_dim, batch_first=True, bidirectional=True)            )            out_dim = hidden_dim * 2            # Project only when the shapes do not already match            self.projections.append(                nn.Identity() if dim == out_dim else nn.Linear(dim, out_dim)            )            dim = out_dim        self.dropout = nn.Dropout(dropout)    def forward(self, x):        for lstm, proj in zip(self.layers, self.projections):            out, _ = lstm(x)            x = self.dropout(out) + proj(x)     # residual        return x

Getting the output shapes right

This is where most bidirectional code breaks, and the failure is silent — no exception, just worse accuracy.

Python
lstm = nn.LSTM(100, 128, num_layers=2, bidirectional=True, batch_first=True)out, (h_n, c_n) = lstm(x)          # x: (32, 50, 100)out.shape    # (32, 50, 256)   -> 2 * hidden_dim, per timesteph_n.shape    # (4, 32, 128)    -> num_layers * num_directions, batch, hidden

The h_n tensor packs directions and layers into its first dimension, ordered as [layer0_fwd, layer0_bwd, layer1_fwd, layer1_bwd]. So:

Python
# WRONG - this is only the backward direction of the last layersummary = h_n[-1]                      # (32, 128)# RIGHT - concatenate both directions of the last layerh = h_n.view(num_layers, 2, batch, hidden)summary = torch.cat([h[-1, 0], h[-1, 1]], dim=1)   # (32, 256)

Take h_n[-1] and you have thrown away the forward pass entirely, halving your model's input. It still trains. It still reports a number. It is just quietly worse than it should be, and nothing tells you.

There is a second, subtler trap. In a bidirectional model the backward pass starts at the end of the padded sequence. Without packing, it begins by processing a long run of pad tokens before it reaches any real content, and its state at the real positions has already been polluted. Packing is even more important here than in a unidirectional model.

Python
packed = nn.utils.rnn.pack_padded_sequence(    emb, lengths.cpu(), batch_first=True, enforce_sorted=False)packed_out, (h_n, c_n) = lstm(packed)out, _ = nn.utils.rnn.pad_packed_sequence(packed_out, batch_first=True)

Attention: choosing which positions matter

A stacked BiLSTM gives you a rich vector at every position. For classification you still have to collapse those into one vector. The obvious choices are weak:

PoolingHowWeakness
Last stateTake hnh_nBiased towards the end of the sequence
MeanAverage all statesOne decisive word is diluted by 200 neutral ones
MaxElement-wise maximumSurprisingly strong, but ignores how many positions agreed
AttentionLearned weighted averageMore parameters; weights can be over-read

Attention computes a relevance score for every position, turns the scores into weights with a softmax, and takes the weighted average.

  1. Score each position: et=score(q,ht)e_t = \text{score}(q, h_t), where qq is a query vector.
  2. Normalise: αt=exp⁡(et)∑kexp⁡(ek)\alpha_t = \dfrac{\exp(e_t)}{\sum_{k} \exp(e_k)}, so weights are positive and sum to 1.
  3. Combine: c=∑tαthtc = \sum_t \alpha_t h_t.

Two scoring functions cover almost all uses:

Additive (Bahdanau)Dot-product (Luong)
Formulav⊤tanh⁡(W1q+W2ht)v^\top \tanh(W_1 q + W_2 h_t)q⊤htq^\top h_t or q⊤Whtq^\top W h_t
ParametersTwo matrices and a vectorNone, or one matrix
Requires matching dimsNoYes (for the plain form)
SpeedSlower — an MLP per pairFast — a single matrix multiply
Best whenSmall models, mismatched dimensionsLarge models, dimensions already aligned

The softmax, with numbers

Suppose a model scores five positions in "the film was not good":

Text
token:   the    film   was    not    goodscore:   0.10   0.80   0.05   2.40   2.10exp:     1.105  2.226  1.051  11.02  8.166      sum = 23.57alpha:   0.047  0.094  0.045  0.468  0.347

Nearly 82% of the weight lands on not and good together. The exponential is what makes attention decisive: a score gap of 2.3 between not and was becomes a weight ratio of over 10 to 1. Small differences in score become large differences in influence.

Python
class AdditiveAttention(nn.Module):    def __init__(self, hidden_dim, attn_dim=128):        super().__init__()        self.W = nn.Linear(hidden_dim, attn_dim, bias=False)        self.v = nn.Linear(attn_dim, 1, bias=False)    def forward(self, states, mask):        # states: (batch, seq, hidden)   mask: (batch, seq) True where real        scores = self.v(torch.tanh(self.W(states))).squeeze(-1)   # (batch, seq)        scores = scores.masked_fill(~mask, float("-inf"))        alpha = torch.softmax(scores, dim=1)        context = torch.bmm(alpha.unsqueeze(1), states).squeeze(1)        return context, alpha

The masked_fill line is mandatory. Padding positions produce scores like anything else, and escoree^{\text{score}} for a pad position is a positive number that takes probability mass away from real tokens. Feeding −∞-\infty makes e−∞=0e^{-\infty} = 0, removing them exactly.

A word of caution on interpretation. Attention weights are frequently presented as an explanation of the model's decision, and they are genuinely useful for debugging — if all the mass sits on <pad> or on the first token, something is broken. But a high weight means "this state contributed heavily to the pooled vector", not "this word caused the prediction". Information from a word can also reach the prediction through the recurrent state of a neighbouring position, without that word ever receiving a high weight.

Putting it together

Python
class BiLSTMAttentionClassifier(nn.Module):    def __init__(self, vocab_size, embed_dim=100, hidden_dim=128,                 num_layers=2, num_classes=2, dropout=0.3, pad_idx=0):        super().__init__()        self.pad_idx = pad_idx        self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=pad_idx)        self.lstm = nn.LSTM(            embed_dim, hidden_dim, num_layers=num_layers,            bidirectional=True, batch_first=True,            dropout=dropout if num_layers > 1 else 0.0,        )        self.attn = AdditiveAttention(hidden_dim * 2)        self.dropout = nn.Dropout(dropout)        self.fc = nn.Linear(hidden_dim * 2, num_classes)   # note the * 2    def forward(self, x, lengths):        mask = x != self.pad_idx        emb = self.dropout(self.embedding(x))        packed = nn.utils.rnn.pack_padded_sequence(            emb, lengths.cpu(), batch_first=True, enforce_sorted=False        )        packed_out, _ = self.lstm(packed)        out, _ = nn.utils.rnn.pad_packed_sequence(            packed_out, batch_first=True, total_length=x.size(1)        )        context, alpha = self.attn(out, mask)        return self.fc(self.dropout(context)), alpha

Two details earn their place. total_length=x.size(1) forces the unpacked output back to the original padded length, so it still lines up with mask — without it, PyTorch trims to the longest sequence in the batch and the shapes silently diverge. And the classifier input is hidden_dim * 2, because the states are concatenations of two directions.

Sequence labelling with the same backbone

For tagging, drop the pooling and put a classifier on every position:

Python
class BiLSTMTagger(nn.Module):    def __init__(self, vocab_size, num_tags, embed_dim=100,                 hidden_dim=128, num_layers=2, dropout=0.3):        super().__init__()        self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=0)        self.lstm = nn.LSTM(embed_dim, hidden_dim, num_layers=num_layers,                            bidirectional=True, batch_first=True,                            dropout=dropout if num_layers > 1 else 0.0)        self.fc = nn.Linear(hidden_dim * 2, num_tags)    def forward(self, x):        out, _ = self.lstm(self.embedding(x))        return self.fc(out)                # (batch, seq, num_tags)criterion = nn.CrossEntropyLoss(ignore_index=-100)loss = criterion(logits.view(-1, num_tags), tags.view(-1))

Set padded target positions to -100 so they are excluded from the loss. A tagger that counts padding in its loss will happily reach 97% "accuracy" by predicting the pad tag, while being useless on real tokens.

One structural weakness remains in this design: each position's tag is predicted independently. The model has no way to express that a tag sequence must be internally consistent — a continuation tag cannot follow "outside", and an entity cannot begin with a continuation. A conditional random field layer on top scores whole tag sequences rather than individual positions and enforces those constraints, typically adding a point or two of entity-level F1 on tagging tasks.

Choosing a configuration

Your taskDirectionLayersPooling
Sentence classification, short inputsBidirectional1Concatenated final states
Document classification, long inputsBidirectional2Attention
Named entity recognition, POS taggingBidirectional2None — per-token output
Text generation, next-word predictionUnidirectional only2–3None
Streaming or real-time inputUnidirectional only1–2Running state
Small dataset (< 5,000 examples)Bidirectional1Mean or max

A sensible build order, one change at a time so you can attribute every gain: start with a single-layer unidirectional LSTM and record the score. Make it bidirectional — expect the largest single jump, especially on tagging. Add attention pooling if the task involves long inputs. Add a second layer last, and keep it only if it earns its place on validation. Everything after that is regularisation.

The thing to carry away is that these are two independent axes, not a ladder. Bidirectionality changes what the model can see; depth changes what it can compose. Bidirectionality is usually the bigger win and is free apart from a doubling of compute — but it is also the one with a hard constraint attached, and if your system will ever have to produce output before the input is complete, no accuracy number justifies using it.