Skip to content

Commit 2587da0

Browse files
authored
Merge pull request #5 from SyntaxSpirits/feature/advanced-architectures
feature: Bidirectional LSTM
2 parents b71112c + 85b2815 commit 2587da0

16 files changed

Lines changed: 928 additions & 92 deletions

CHANGELOG.md

Lines changed: 26 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,32 @@ All notable changes to this project will be documented in this file.
55
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/),
66
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
77

8+
## [0.3.0] - 2025-06-15
9+
10+
### Added
11+
- **Bidirectional LSTM Networks**: Complete implementation of BiLSTM with flexible output combination modes
12+
- **Multiple Combine Modes**: Support for concatenation, sum, and average combination of forward/backward outputs
13+
- **Multi-layer BiLSTM**: Stacked bidirectional layers with proper input/output size handling
14+
- **BiLSTM Training Support**: Forward pass with caching for efficient gradient computation
15+
- **Dropout Compatibility**: Full support for input, recurrent, output dropout and zoneout in BiLSTM
16+
- **Examples and Documentation**:
17+
- Comprehensive BiLSTM demonstration example
18+
- Text classification comparison example (BiLSTM vs unidirectional LSTM)
19+
- Updated documentation with usage examples
20+
- **API Integration**: BiLSTM seamlessly integrated with existing training and optimization systems
21+
22+
### Technical Details
23+
- Implemented proper bidirectional sequence processing with forward and backward passes
24+
- Added efficient multi-layer processing with correct input dimension handling
25+
- Created comprehensive test suite for BiLSTM functionality
26+
- Maintained compatibility with existing LSTM training infrastructure
27+
28+
### Benefits
29+
- Better context understanding for sequence modeling tasks
30+
- Improved performance on sequence labeling and classification
31+
- Access to both past and future context for each time step
32+
- Flexible output combination strategies for different use cases
33+
834
## [0.2.0] - 2025-06-09
935

1036
### Added

Cargo.toml

Lines changed: 9 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,6 @@
11
[package]
22
name = "rust-lstm"
3-
version = "0.2.0"
3+
version = "0.3.0"
44
authors = ["Alex Kholodniak <alexandrkholodniak@gmail.com>"]
55
edition = "2021"
66
description = "A complete LSTM neural network library with training capabilities, multiple optimizers, and peephole variants."
@@ -51,3 +51,11 @@ path = "examples/real_data_example.rs"
5151
[[example]]
5252
name = "model_inspection"
5353
path = "examples/model_inspection.rs"
54+
55+
[[example]]
56+
name = "bilstm_example"
57+
path = "examples/bilstm_example.rs"
58+
59+
[[example]]
60+
name = "text_classification_bilstm"
61+
path = "examples/text_classification_bilstm.rs"

README.md

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@ A comprehensive LSTM (Long Short-Term Memory) neural network library implemented
66

77
- **LSTM cell implementation** with forward and backward propagation
88
- **Peephole LSTM variant** for enhanced performance
9+
- **Bidirectional LSTM networks** with flexible output combination modes
910
- **Multi-layer LSTM networks** with configurable architecture
1011
- **Complete training system** with backpropagation through time (BPTT)
1112
- **Multiple optimizers**: SGD, Adam, RMSprop
@@ -220,6 +221,42 @@ let mut cell = PeepholeLSTMCell::new(input_size, hidden_size);
220221
let (h_t, c_t) = cell.forward(&input, &h_prev, &c_prev);
221222
```
222223

224+
#### Bidirectional LSTM
225+
226+
```rust
227+
use rust_lstm::layers::bilstm_network::{BiLSTMNetwork, CombineMode};
228+
229+
// Create BiLSTM with concatenated outputs (most common)
230+
let mut bilstm = BiLSTMNetwork::new_concat(input_size, hidden_size, num_layers);
231+
232+
// Create BiLSTM with different combine modes
233+
let bilstm_sum = BiLSTMNetwork::new_sum(input_size, hidden_size, num_layers);
234+
let bilstm_avg = BiLSTMNetwork::new_average(input_size, hidden_size, num_layers);
235+
236+
// Or specify combine mode explicitly
237+
let bilstm_custom = BiLSTMNetwork::new(input_size, hidden_size, num_layers, CombineMode::Concat);
238+
239+
// Process a sequence (captures both past and future context)
240+
let sequence = vec![
241+
Array2::from_shape_vec((input_size, 1), vec![0.1, 0.2]).unwrap(),
242+
Array2::from_shape_vec((input_size, 1), vec![0.3, 0.4]).unwrap(),
243+
Array2::from_shape_vec((input_size, 1), vec![0.5, 0.6]).unwrap(),
244+
];
245+
246+
let outputs = bilstm.forward_sequence(&sequence);
247+
248+
// BiLSTM with dropout
249+
let mut bilstm = BiLSTMNetwork::new_concat(input_size, hidden_size, num_layers)
250+
.with_input_dropout(0.2, true) // Variational input dropout
251+
.with_recurrent_dropout(0.3, true) // Variational recurrent dropout
252+
.with_output_dropout(0.1); // Standard output dropout
253+
254+
// Output size depends on combine mode:
255+
// - Concat: 2 * hidden_size
256+
// - Sum/Average: hidden_size
257+
println!("Output size: {}", bilstm.output_size());
258+
```
259+
223260
To run this example, save it as main.rs, and run:
224261

225262
```bash
@@ -233,6 +270,7 @@ The library includes several examples demonstrating different use cases:
233270
- `basic_usage.rs` - Simple forward pass example
234271
- `training_example.rs` - Complete training workflow with multiple optimizers
235272
- `dropout_example.rs` - Comprehensive dropout regularization demo
273+
- `bilstm_example.rs` - Bidirectional LSTM demonstration with different combine modes
236274
- `time_series_prediction.rs` - Time series forecasting
237275
- `text_generation_advanced.rs` - Character-level text generation
238276
- `multi_layer_lstm.rs` - Multi-layer network usage
@@ -246,6 +284,7 @@ Run examples with:
246284
```bash
247285
cargo run --example training_example
248286
cargo run --example dropout_example
287+
cargo run --example bilstm_example
249288
cargo run --example time_series_prediction
250289
cargo run --example stock_prediction
251290
cargo run --example weather_prediction

examples/bilstm_example.rs

Lines changed: 202 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,202 @@
1+
use ndarray::{Array2, arr2};
2+
use rust_lstm::layers::bilstm_network::{BiLSTMNetwork, CombineMode};
3+
use rust_lstm::models::lstm_network::LSTMNetwork;
4+
5+
/// Generate a simple sequence that benefits from bidirectional processing
6+
fn generate_bidirectional_data() -> Vec<Array2<f64>> {
7+
let sequence_length = 10;
8+
let mut sequence = Vec::new();
9+
10+
for t in 0..sequence_length {
11+
let t_f = t as f64 * 0.5;
12+
let current = t_f.sin();
13+
let future = if t < sequence_length - 1 { (t_f + 0.5).cos() * 0.5 } else { 0.0 };
14+
let past = if t > 0 { (t_f - 0.5).sin() * 0.3 } else { 0.0 };
15+
16+
let value = current + future + past;
17+
sequence.push(arr2(&[[value]]));
18+
}
19+
20+
sequence
21+
}
22+
23+
/// Demonstrate basic BiLSTM functionality
24+
fn demo_basic_bilstm() {
25+
println!("=== Basic BiLSTM Demonstration ===");
26+
27+
let mut bilstm = BiLSTMNetwork::new_concat(1, 4, 1);
28+
let sequence = generate_bidirectional_data();
29+
30+
println!("Input sequence length: {}", sequence.len());
31+
println!("BiLSTM hidden size: {}", bilstm.hidden_size);
32+
println!("BiLSTM output size: {}", bilstm.output_size());
33+
34+
let outputs = bilstm.forward_sequence(&sequence);
35+
36+
println!("Output shapes:");
37+
for (i, output) in outputs.iter().enumerate() {
38+
println!(" Time step {}: {:?}", i, output.shape());
39+
}
40+
41+
println!("Sample output values (first 3 time steps):");
42+
for (i, output) in outputs.iter().take(3).enumerate() {
43+
println!(" t={}: [{:.4}, {:.4}, {:.4}, ...]",
44+
i, output[[0,0]], output[[1,0]], output[[2,0]]);
45+
}
46+
}
47+
48+
/// Compare different combine modes
49+
fn demo_combine_modes() {
50+
println!("\n=== BiLSTM Combine Modes Comparison ===");
51+
52+
let sequence = generate_bidirectional_data();
53+
54+
// Test different combine modes
55+
let modes = vec![
56+
("Concatenation", CombineMode::Concat),
57+
("Sum", CombineMode::Sum),
58+
("Average", CombineMode::Average),
59+
];
60+
61+
for (name, mode) in modes {
62+
let mut bilstm = BiLSTMNetwork::new(1, 3, 1, mode);
63+
let outputs = bilstm.forward_sequence(&sequence);
64+
65+
println!("{} mode:", name);
66+
println!(" Output size: {}", bilstm.output_size());
67+
println!(" First output shape: {:?}", outputs[0].shape());
68+
println!(" Sample values: [{:.4}, {:.4}]",
69+
outputs[0][[0,0]],
70+
if outputs[0].nrows() > 1 { outputs[0][[1,0]] } else { 0.0 });
71+
}
72+
}
73+
74+
/// Compare BiLSTM vs unidirectional LSTM performance
75+
fn demo_bilstm_vs_lstm() {
76+
println!("\n=== BiLSTM vs Unidirectional LSTM Comparison ===");
77+
78+
let sequence = generate_bidirectional_data();
79+
80+
// Unidirectional LSTM
81+
let mut lstm = LSTMNetwork::new(1, 4, 1);
82+
let mut hx = Array2::zeros((4, 1));
83+
let mut cx = Array2::zeros((4, 1));
84+
85+
let mut lstm_outputs = Vec::new();
86+
for input in &sequence {
87+
let (new_hx, new_cx) = lstm.forward(input, &hx, &cx);
88+
lstm_outputs.push(new_hx.clone());
89+
hx = new_hx;
90+
cx = new_cx;
91+
}
92+
93+
// Bidirectional LSTM (with same total parameters approximately)
94+
let mut bilstm = BiLSTMNetwork::new_concat(1, 2, 1); // 2*2=4 total hidden units
95+
let bilstm_outputs = bilstm.forward_sequence(&sequence);
96+
97+
println!("Unidirectional LSTM:");
98+
println!(" Hidden size: 4");
99+
println!(" Output size: 4");
100+
println!(" Sample output: [{:.4}, {:.4}, {:.4}, {:.4}]",
101+
lstm_outputs[0][[0,0]], lstm_outputs[0][[1,0]],
102+
lstm_outputs[0][[2,0]], lstm_outputs[0][[3,0]]);
103+
104+
println!("Bidirectional LSTM:");
105+
println!(" Hidden size per direction: 2");
106+
println!(" Total output size: 4");
107+
println!(" Sample output: [{:.4}, {:.4}, {:.4}, {:.4}]",
108+
bilstm_outputs[0][[0,0]], bilstm_outputs[0][[1,0]],
109+
bilstm_outputs[0][[2,0]], bilstm_outputs[0][[3,0]]);
110+
111+
// Demonstrate that BiLSTM has access to future context
112+
println!("\nContext Analysis:");
113+
println!(" LSTM processes left-to-right only");
114+
println!(" BiLSTM processes both directions and combines information");
115+
println!(" This allows BiLSTM to use future context for current predictions");
116+
}
117+
118+
/// Demonstrate multi-layer BiLSTM
119+
fn demo_multilayer_bilstm() {
120+
println!("\n=== Multi-layer BiLSTM ===");
121+
122+
let sequence = generate_bidirectional_data();
123+
124+
for num_layers in 1..=3 {
125+
let mut bilstm = BiLSTMNetwork::new_concat(1, 3, num_layers);
126+
let outputs = bilstm.forward_sequence(&sequence);
127+
128+
println!("{}-layer BiLSTM:", num_layers);
129+
println!(" Total parameters (approx): {}",
130+
num_layers * 2 * (3 * 4 * (if num_layers == 1 { 1 } else { 6 }) + 4 * 3));
131+
println!(" Output shape: {:?}", outputs[0].shape());
132+
println!(" Sample output magnitude: {:.4}",
133+
outputs[0].iter().map(|&x| x.abs()).sum::<f64>() / outputs[0].len() as f64);
134+
}
135+
}
136+
137+
/// Demonstrate BiLSTM with dropout
138+
fn demo_bilstm_with_dropout() {
139+
println!("\n=== BiLSTM with Dropout ===");
140+
141+
let sequence = generate_bidirectional_data();
142+
143+
let mut bilstm = BiLSTMNetwork::new_concat(1, 4, 2)
144+
.with_input_dropout(0.2, true) // 20% variational input dropout
145+
.with_recurrent_dropout(0.3, true) // 30% variational recurrent dropout
146+
.with_output_dropout(0.1); // 10% output dropout
147+
148+
// Training mode (dropout active)
149+
bilstm.train();
150+
let train_outputs = bilstm.forward_sequence(&sequence);
151+
152+
// Evaluation mode (dropout inactive)
153+
bilstm.eval();
154+
let eval_outputs = bilstm.forward_sequence(&sequence);
155+
156+
println!("Training mode (with dropout):");
157+
println!(" Sample output: [{:.4}, {:.4}, {:.4}]",
158+
train_outputs[0][[0,0]], train_outputs[0][[1,0]], train_outputs[0][[2,0]]);
159+
160+
println!("Evaluation mode (no dropout):");
161+
println!(" Sample output: [{:.4}, {:.4}, {:.4}]",
162+
eval_outputs[0][[0,0]], eval_outputs[0][[1,0]], eval_outputs[0][[2,0]]);
163+
164+
println!("Dropout correctly affects training vs evaluation outputs");
165+
}
166+
167+
/// Demonstrate sequence processing with caching
168+
fn demo_bilstm_with_caching() {
169+
println!("\n=== BiLSTM with Caching (for Training) ===");
170+
171+
let sequence = generate_bidirectional_data();
172+
let mut bilstm = BiLSTMNetwork::new_concat(1, 3, 1);
173+
174+
let (outputs, cache) = bilstm.forward_sequence_with_cache(&sequence);
175+
176+
println!("Forward pass with caching:");
177+
println!(" Sequence length: {}", sequence.len());
178+
println!(" Number of outputs: {}", outputs.len());
179+
println!(" Forward caches: {}", cache.forward_caches.len());
180+
println!(" Backward caches: {}", cache.backward_caches.len());
181+
println!(" Cache enables efficient backpropagation for training");
182+
}
183+
184+
fn main() {
185+
println!("🔄 Bidirectional LSTM Demonstration");
186+
println!("=====================================");
187+
188+
demo_basic_bilstm();
189+
demo_combine_modes();
190+
demo_bilstm_vs_lstm();
191+
demo_multilayer_bilstm();
192+
demo_bilstm_with_dropout();
193+
demo_bilstm_with_caching();
194+
195+
println!("\n✅ BiLSTM demonstration completed!");
196+
println!("\nKey Benefits of Bidirectional LSTM:");
197+
println!("• Captures both past and future context");
198+
println!("• Better for tasks where full sequence is available");
199+
println!("• Improved performance on sequence labeling tasks");
200+
println!("• Flexible output combination modes");
201+
println!("• Compatible with existing dropout and training systems");
202+
}

examples/dropout_example.rs

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -28,9 +28,9 @@ fn demonstrate_basic_dropout() {
2828

2929
// Create network with standard dropout
3030
let mut network = LSTMNetwork::new(input_size, hidden_size, num_layers)
31-
.with_input_dropout(0.2, false) // 20% input dropout (non-variational)
32-
.with_recurrent_dropout(0.3, false) // 30% recurrent dropout (non-variational)
33-
.with_output_dropout(0.1); // 10% output dropout
31+
.with_input_dropout(0.2, false)
32+
.with_recurrent_dropout(0.3, false)
33+
.with_output_dropout(0.1);
3434

3535
let input = arr2(&[[1.0], [0.5], [-0.2], [0.8]]);
3636
let hx = Array2::zeros((hidden_size, 1));
@@ -65,8 +65,8 @@ fn demonstrate_variational_dropout() {
6565

6666
// Create network with variational dropout (same mask across time steps)
6767
let mut network = LSTMNetwork::new(input_size, hidden_size, num_layers)
68-
.with_input_dropout(0.25, true) // Variational input dropout
69-
.with_recurrent_dropout(0.2, true); // Variational recurrent dropout
68+
.with_input_dropout(0.25, true)
69+
.with_recurrent_dropout(0.2, true);
7070

7171
let sequence = vec![
7272
arr2(&[[1.0], [0.0], [0.5]]),

0 commit comments

Comments
 (0)