Skip to content

Recipe: RNN vs. LSTM vs. GRU on a Parity Task

What you'll build: three recurrent modules (RNNModule, LSTMModule, GRUModule) trained side by side on the same synthetic running-parity (cumulative XOR) task. The comparison shows why gated recurrent units exist.

CMake target: sequence_models_recipe (examples/recipes/sequence_models_rnn_lstm_gru.cpp).

Run it: ./build/sequence_models_recipe (Windows: build\Release\sequence_models_recipe.exe).

Code

RNNModule rnn(1, 1, &backend);
// ... seeded weight init ...

LSTMModule lstm(1, 1, &backend);
// ... seeded weight init ...

GRUModule gru(1, 1, &backend);
// ... seeded weight init ...

// train_and_report(): the same loop runs for rnn, lstm and gru
AdamOptimizer optimizer(0.05f, &backend);
MSELoss loss(&backend);
float final_loss = 0.0f;
for (int epoch = 0; epoch <= 2000; ++epoch) {
    optimizer.zero_grad(module);
    Tensor prediction = module.forward(input);   // input: (6, 10, 1) binary sequences
    final_loss = loss.forward(prediction, target);
    if (epoch == 2000) break;                    // last pass only measures the loss
    Tensor grad = loss.backward();
    (void)module.backward(grad);
    optimizer.step(module);
}

Full source: examples/recipes/sequence_models_rnn_lstm_gru.cpp.

Expected output

Sequence-model recipe -- running parity of 6 seeded sequences (length 10)

RNNModule  | final loss 0.241627 | timestep accuracy  32/60 (53.3%)
LSTMModule | final loss 0.013926 | timestep accuracy  60/60 (100.0%)
GRUModule  | final loss 0.000020 | timestep accuracy  60/60 (100.0%)
...

What's happening

The task is h_t = XOR(h_{t-1}, x_t). A vanilla Elman cell computes tanh(w_x*x_t + w_h*h_{t-1} + b), which is monotone in each argument. A monotone function cannot represent XOR (see examples/sequence_model_demo.cpp for the full argument). So RNNModule settles on the best monotone fit: predict the current bit and ignore history. It lands on loss 0.2416 regardless of learning rate or seed.

LSTMModule and GRUModule solve the task exactly. Their multiplicative gates let the carried state flip sign depending on the current input, which a plain tanh unit can't do. The RNN's failure is a limit of what it can represent, not an optimization failure.

See also: Deep Learning Modules and Layers.