Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 26 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,32 @@ All notable changes to this project will be documented in this file.
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/),
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).

## [0.3.0] - 2025-06-15

### Added
- **Bidirectional LSTM Networks**: Complete implementation of BiLSTM with flexible output combination modes
- **Multiple Combine Modes**: Support for concatenation, sum, and average combination of forward/backward outputs
- **Multi-layer BiLSTM**: Stacked bidirectional layers with proper input/output size handling
- **BiLSTM Training Support**: Forward pass with caching for efficient gradient computation
- **Dropout Compatibility**: Full support for input, recurrent, output dropout and zoneout in BiLSTM
- **Examples and Documentation**:
- Comprehensive BiLSTM demonstration example
- Text classification comparison example (BiLSTM vs unidirectional LSTM)
- Updated documentation with usage examples
- **API Integration**: BiLSTM seamlessly integrated with existing training and optimization systems

### Technical Details
- Implemented proper bidirectional sequence processing with forward and backward passes
- Added efficient multi-layer processing with correct input dimension handling
- Created comprehensive test suite for BiLSTM functionality
- Maintained compatibility with existing LSTM training infrastructure

### Benefits
- Better context understanding for sequence modeling tasks
- Improved performance on sequence labeling and classification
- Access to both past and future context for each time step
- Flexible output combination strategies for different use cases

## [0.2.0] - 2025-06-09

### Added
Expand Down
10 changes: 9 additions & 1 deletion Cargo.toml
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
[package]
name = "rust-lstm"
version = "0.2.0"
version = "0.3.0"
authors = ["Alex Kholodniak <alexandrkholodniak@gmail.com>"]
edition = "2021"
description = "A complete LSTM neural network library with training capabilities, multiple optimizers, and peephole variants."
Expand Down Expand Up @@ -51,3 +51,11 @@ path = "examples/real_data_example.rs"
[[example]]
name = "model_inspection"
path = "examples/model_inspection.rs"

[[example]]
name = "bilstm_example"
path = "examples/bilstm_example.rs"

[[example]]
name = "text_classification_bilstm"
path = "examples/text_classification_bilstm.rs"
39 changes: 39 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ A comprehensive LSTM (Long Short-Term Memory) neural network library implemented

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

#### Bidirectional LSTM

```rust
use rust_lstm::layers::bilstm_network::{BiLSTMNetwork, CombineMode};

// Create BiLSTM with concatenated outputs (most common)
let mut bilstm = BiLSTMNetwork::new_concat(input_size, hidden_size, num_layers);

// Create BiLSTM with different combine modes
let bilstm_sum = BiLSTMNetwork::new_sum(input_size, hidden_size, num_layers);
let bilstm_avg = BiLSTMNetwork::new_average(input_size, hidden_size, num_layers);

// Or specify combine mode explicitly
let bilstm_custom = BiLSTMNetwork::new(input_size, hidden_size, num_layers, CombineMode::Concat);

// Process a sequence (captures both past and future context)
let sequence = vec![
Array2::from_shape_vec((input_size, 1), vec![0.1, 0.2]).unwrap(),
Array2::from_shape_vec((input_size, 1), vec![0.3, 0.4]).unwrap(),
Array2::from_shape_vec((input_size, 1), vec![0.5, 0.6]).unwrap(),
];

let outputs = bilstm.forward_sequence(&sequence);

// BiLSTM with dropout
let mut bilstm = BiLSTMNetwork::new_concat(input_size, hidden_size, num_layers)
.with_input_dropout(0.2, true) // Variational input dropout
.with_recurrent_dropout(0.3, true) // Variational recurrent dropout
.with_output_dropout(0.1); // Standard output dropout

// Output size depends on combine mode:
// - Concat: 2 * hidden_size
// - Sum/Average: hidden_size
println!("Output size: {}", bilstm.output_size());
```

To run this example, save it as main.rs, and run:

```bash
Expand All @@ -233,6 +270,7 @@ The library includes several examples demonstrating different use cases:
- `basic_usage.rs` - Simple forward pass example
- `training_example.rs` - Complete training workflow with multiple optimizers
- `dropout_example.rs` - Comprehensive dropout regularization demo
- `bilstm_example.rs` - Bidirectional LSTM demonstration with different combine modes
- `time_series_prediction.rs` - Time series forecasting
- `text_generation_advanced.rs` - Character-level text generation
- `multi_layer_lstm.rs` - Multi-layer network usage
Expand All @@ -246,6 +284,7 @@ Run examples with:
```bash
cargo run --example training_example
cargo run --example dropout_example
cargo run --example bilstm_example
cargo run --example time_series_prediction
cargo run --example stock_prediction
cargo run --example weather_prediction
Expand Down
202 changes: 202 additions & 0 deletions examples/bilstm_example.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,202 @@
use ndarray::{Array2, arr2};
use rust_lstm::layers::bilstm_network::{BiLSTMNetwork, CombineMode};
use rust_lstm::models::lstm_network::LSTMNetwork;

/// Generate a simple sequence that benefits from bidirectional processing
fn generate_bidirectional_data() -> Vec<Array2<f64>> {
let sequence_length = 10;
let mut sequence = Vec::new();

for t in 0..sequence_length {
let t_f = t as f64 * 0.5;
let current = t_f.sin();
let future = if t < sequence_length - 1 { (t_f + 0.5).cos() * 0.5 } else { 0.0 };
let past = if t > 0 { (t_f - 0.5).sin() * 0.3 } else { 0.0 };

let value = current + future + past;
sequence.push(arr2(&[[value]]));
}

sequence
}

/// Demonstrate basic BiLSTM functionality
fn demo_basic_bilstm() {
println!("=== Basic BiLSTM Demonstration ===");

let mut bilstm = BiLSTMNetwork::new_concat(1, 4, 1);
let sequence = generate_bidirectional_data();

println!("Input sequence length: {}", sequence.len());
println!("BiLSTM hidden size: {}", bilstm.hidden_size);
println!("BiLSTM output size: {}", bilstm.output_size());

let outputs = bilstm.forward_sequence(&sequence);

println!("Output shapes:");
for (i, output) in outputs.iter().enumerate() {
println!(" Time step {}: {:?}", i, output.shape());
}

println!("Sample output values (first 3 time steps):");
for (i, output) in outputs.iter().take(3).enumerate() {
println!(" t={}: [{:.4}, {:.4}, {:.4}, ...]",
i, output[[0,0]], output[[1,0]], output[[2,0]]);
}
}

/// Compare different combine modes
fn demo_combine_modes() {
println!("\n=== BiLSTM Combine Modes Comparison ===");

let sequence = generate_bidirectional_data();

// Test different combine modes
let modes = vec![
("Concatenation", CombineMode::Concat),
("Sum", CombineMode::Sum),
("Average", CombineMode::Average),
];

for (name, mode) in modes {
let mut bilstm = BiLSTMNetwork::new(1, 3, 1, mode);
let outputs = bilstm.forward_sequence(&sequence);

println!("{} mode:", name);
println!(" Output size: {}", bilstm.output_size());
println!(" First output shape: {:?}", outputs[0].shape());
println!(" Sample values: [{:.4}, {:.4}]",
outputs[0][[0,0]],
if outputs[0].nrows() > 1 { outputs[0][[1,0]] } else { 0.0 });
}
}

/// Compare BiLSTM vs unidirectional LSTM performance
fn demo_bilstm_vs_lstm() {
println!("\n=== BiLSTM vs Unidirectional LSTM Comparison ===");

let sequence = generate_bidirectional_data();

// Unidirectional LSTM
let mut lstm = LSTMNetwork::new(1, 4, 1);
let mut hx = Array2::zeros((4, 1));
let mut cx = Array2::zeros((4, 1));

let mut lstm_outputs = Vec::new();
for input in &sequence {
let (new_hx, new_cx) = lstm.forward(input, &hx, &cx);
lstm_outputs.push(new_hx.clone());
hx = new_hx;
cx = new_cx;
}

// Bidirectional LSTM (with same total parameters approximately)
let mut bilstm = BiLSTMNetwork::new_concat(1, 2, 1); // 2*2=4 total hidden units
let bilstm_outputs = bilstm.forward_sequence(&sequence);

println!("Unidirectional LSTM:");
println!(" Hidden size: 4");
println!(" Output size: 4");
println!(" Sample output: [{:.4}, {:.4}, {:.4}, {:.4}]",
lstm_outputs[0][[0,0]], lstm_outputs[0][[1,0]],
lstm_outputs[0][[2,0]], lstm_outputs[0][[3,0]]);

println!("Bidirectional LSTM:");
println!(" Hidden size per direction: 2");
println!(" Total output size: 4");
println!(" Sample output: [{:.4}, {:.4}, {:.4}, {:.4}]",
bilstm_outputs[0][[0,0]], bilstm_outputs[0][[1,0]],
bilstm_outputs[0][[2,0]], bilstm_outputs[0][[3,0]]);

// Demonstrate that BiLSTM has access to future context
println!("\nContext Analysis:");
println!(" LSTM processes left-to-right only");
println!(" BiLSTM processes both directions and combines information");
println!(" This allows BiLSTM to use future context for current predictions");
}

/// Demonstrate multi-layer BiLSTM
fn demo_multilayer_bilstm() {
println!("\n=== Multi-layer BiLSTM ===");

let sequence = generate_bidirectional_data();

for num_layers in 1..=3 {
let mut bilstm = BiLSTMNetwork::new_concat(1, 3, num_layers);
let outputs = bilstm.forward_sequence(&sequence);

println!("{}-layer BiLSTM:", num_layers);
println!(" Total parameters (approx): {}",
num_layers * 2 * (3 * 4 * (if num_layers == 1 { 1 } else { 6 }) + 4 * 3));
println!(" Output shape: {:?}", outputs[0].shape());
println!(" Sample output magnitude: {:.4}",
outputs[0].iter().map(|&x| x.abs()).sum::<f64>() / outputs[0].len() as f64);
}
}

/// Demonstrate BiLSTM with dropout
fn demo_bilstm_with_dropout() {
println!("\n=== BiLSTM with Dropout ===");

let sequence = generate_bidirectional_data();

let mut bilstm = BiLSTMNetwork::new_concat(1, 4, 2)
.with_input_dropout(0.2, true) // 20% variational input dropout
.with_recurrent_dropout(0.3, true) // 30% variational recurrent dropout
.with_output_dropout(0.1); // 10% output dropout

// Training mode (dropout active)
bilstm.train();
let train_outputs = bilstm.forward_sequence(&sequence);

// Evaluation mode (dropout inactive)
bilstm.eval();
let eval_outputs = bilstm.forward_sequence(&sequence);

println!("Training mode (with dropout):");
println!(" Sample output: [{:.4}, {:.4}, {:.4}]",
train_outputs[0][[0,0]], train_outputs[0][[1,0]], train_outputs[0][[2,0]]);

println!("Evaluation mode (no dropout):");
println!(" Sample output: [{:.4}, {:.4}, {:.4}]",
eval_outputs[0][[0,0]], eval_outputs[0][[1,0]], eval_outputs[0][[2,0]]);

println!("Dropout correctly affects training vs evaluation outputs");
}

/// Demonstrate sequence processing with caching
fn demo_bilstm_with_caching() {
println!("\n=== BiLSTM with Caching (for Training) ===");

let sequence = generate_bidirectional_data();
let mut bilstm = BiLSTMNetwork::new_concat(1, 3, 1);

let (outputs, cache) = bilstm.forward_sequence_with_cache(&sequence);

println!("Forward pass with caching:");
println!(" Sequence length: {}", sequence.len());
println!(" Number of outputs: {}", outputs.len());
println!(" Forward caches: {}", cache.forward_caches.len());
println!(" Backward caches: {}", cache.backward_caches.len());
println!(" Cache enables efficient backpropagation for training");
}

fn main() {
println!("🔄 Bidirectional LSTM Demonstration");
println!("=====================================");

demo_basic_bilstm();
demo_combine_modes();
demo_bilstm_vs_lstm();
demo_multilayer_bilstm();
demo_bilstm_with_dropout();
demo_bilstm_with_caching();

println!("\n✅ BiLSTM demonstration completed!");
println!("\nKey Benefits of Bidirectional LSTM:");
println!("• Captures both past and future context");
println!("• Better for tasks where full sequence is available");
println!("• Improved performance on sequence labeling tasks");
println!("• Flexible output combination modes");
println!("• Compatible with existing dropout and training systems");
}
10 changes: 5 additions & 5 deletions examples/dropout_example.rs
Original file line number Diff line number Diff line change
Expand Up @@ -28,9 +28,9 @@ fn demonstrate_basic_dropout() {

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

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

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

let sequence = vec![
arr2(&[[1.0], [0.0], [0.5]]),
Expand Down
Loading