🚀 Unofficial implementation of A Simple RNN Model for Lightweight, Low-compute and Low-latency Multichannel Speech Enhancement in the Time Domain in PyTorch.
This project implements a lightweight RNN-based model for real-time multichannel speech enhancement in the time domain. The model is designed for low-compute and low-latency applications, making it suitable for edge devices and real-time processing scenarios.
The SimpleRNNModel consists of the following components:
- Input Projection: Linear layer + LayerNorm + PReLU
- Spatial Filters: Learnable spatial filtering mechanism for multichannel processing
- RNN Layers: Stack of LSTM layers with LayerNorm (default: 3 layers)
- Output Projection: Linear layer for enhanced signal reconstruction
n_channels: Number of input channels (default: 8)hidden_dim: Hidden dimension of RNN (default: 256)iW: Input window size in samples (e.g., 32 samples = 2ms at 16kHz)oW: Output window size in samples (e.g., 32 samples = 2ms at 16kHz)S: Stride in samples (default: 16 samples = 1ms at 16kHz)B: Number of RNN layers (default: 3)
from simplernn import SimpleRNNModel
# Initialize model
model = SimpleRNNModel(
n_channels=8, # Number of input channels
hidden_dim=256, # Hidden dimension
iW=32, # Input window size (2ms at 16kHz)
oW=32, # Output window size (2ms at 16kHz)
S=16, # Stride (1ms at 16kHz)
B=3 # Number of RNN layers
)
# Forward pass
# Input shape: [batch_size, n_channels, num_samples] (e.g., [1, 8, 16000])
# Output shape: [batch_size, num_samples] (Single-channel enhanced audio)
enhanced_audio = model(input_audio)The repository includes comprehensive complexity analysis for various model configurations:
python simplernn.pyThis will output FLOPs and parameter counts for different model variants.
Note: All FLOPs (GMacs/MMacs) listed below are calculated for processing 1 second of 16kHz audio (
| Model | Channels | Hidden Dim | Input Window | Output Window | FLOPs | Parameters |
|---|---|---|---|---|---|---|
| L4_C2_H300_a | 2 | 300 | 64 (4ms) | 64 (4ms) | 2.23 GMac | 2.21 M |
| L4_C4_H300_a | 4 | 300 | 64 (4ms) | 64 (4ms) | 2.27 GMac | 2.21 M |
| L4_C8_H300_a | 8 | 300 | 64 (4ms) | 64 (4ms) | 2.35 GMac | 2.21 M |
| L2_C2_H300_a | 2 | 300 | 32 (2ms) | 32 (2ms) | 2.21 GMac | 2.19 M |
| L2_C4_H300_a | 4 | 300 | 32 (2ms) | 32 (2ms) | 2.23 GMac | 2.19 M |
| L2_C8_H300_a | 8 | 300 | 32 (2ms) | 32 (2ms) | 2.27 GMac | 2.19 M |
| L1_C2_H300_a | 2 | 300 | 16 (1ms) | 16 (1ms) | 2.19 GMac | 2.18 M |
| L1_C4_H300_a | 4 | 300 | 16 (1ms) | 16 (1ms) | 2.21 GMac | 2.18 M |
| L1_C8_H300_a | 8 | 300 | 16 (1ms) | 16 (1ms) | 2.23 GMac | 2.18 M |
Approach (b): Fixed input window (iW = 256 samples = 16ms), output window varies with latency (oW = LĂ—16)
| Model | Channels | Hidden Dim | Input Window | Output Window | FLOPs | Parameters |
|---|---|---|---|---|---|---|
| L4_C2_H300_b | 2 | 300 | 256 (16ms) | 64 (4ms) | 2.35 GMac | 2.27 M |
| L4_C4_H300_b | 4 | 300 | 256 (16ms) | 64 (4ms) | 2.50 GMac | 2.27 M |
| L4_C8_H300_b | 8 | 300 | 256 (16ms) | 64 (4ms) | 2.81 GMac | 2.27 M |
| L2_C2_H300_b | 2 | 300 | 256 (16ms) | 32 (2ms) | 2.34 GMac | 2.26 M |
| L2_C4_H300_b | 4 | 300 | 256 (16ms) | 32 (2ms) | 2.50 GMac | 2.26 M |
| L2_C8_H300_b | 8 | 300 | 256 (16ms) | 32 (2ms) | 2.81 GMac | 2.26 M |
| L1_C2_H300_b | 2 | 300 | 256 (16ms) | 16 (1ms) | 2.34 GMac | 2.25 M |
| L1_C4_H300_b | 4 | 300 | 256 (16ms) | 16 (1ms) | 2.49 GMac | 2.25 M |
| L1_C8_H300_b | 8 | 300 | 256 (16ms) | 16 (1ms) | 2.81 GMac | 2.25 M |
Approach (b) with larger hidden dimension (H1024)
| Model | Channels | Hidden Dim | Input Window | Output Window | FLOPs | Parameters |
|---|---|---|---|---|---|---|
| L4_C2_H1024_b | 2 | 1024 | 256 (16ms) | 64 (4ms) | 25.74 GMac | 25.53 M |
| L4_C4_H1024_b | 4 | 1024 | 256 (16ms) | 64 (4ms) | 26.28 GMac | 25.53 M |
| L4_C8_H1024_b | 8 | 1024 | 256 (16ms) | 64 (4ms) | 27.34 GMac | 25.54 M |
| L2_C2_H1024_b | 2 | 1024 | 256 (16ms) | 32 (2ms) | 25.76 GMac | 25.50 M |
| L2_C4_H1024_b | 4 | 1024 | 256 (16ms) | 32 (2ms) | 26.30 GMac | 25.50 M |
| L2_C8_H1024_b | 8 | 1024 | 256 (16ms) | 32 (2ms) | 27.36 GMac | 25.50 M |
| L1_C2_H1024_b | 2 | 1024 | 256 (16ms) | 16 (1ms) | 25.77 GMac | 25.48 M |
| L1_C4_H1024_b | 4 | 1024 | 256 (16ms) | 16 (1ms) | 26.31 GMac | 25.48 M |
| L1_C8_H1024_b | 8 | 1024 | 256 (16ms) | 16 (1ms) | 27.37 GMac | 25.49 M |
Performance of the proposed model for an algorithmic latency of 2 ms with varying widths.
| Model | Channels | Hidden Dim | FLOPs | Parameters |
|---|---|---|---|---|
| H64_C2_a | 2 | 64 | 108.53 MMac | 104.67 k |
| H64_C4_a | 4 | 64 | 113.13 MMac | 104.80 k |
| H64_C8_a | 8 | 64 | 122.34 MMac | 105.06 k |
| H128_C2_a | 2 | 128 | 413.44 MMac | 405.92 k |
| H128_C4_a | 4 | 128 | 422.65 MMac | 406.18 k |
| H128_C8_a | 8 | 128 | 441.06 MMac | 406.69 k |
| H256_C2_a | 2 | 256 | 1.61 GMac | 1.60 M |
| H256_C4_a | 4 | 256 | 1.63 GMac | 1.60 M |
| H256_C8_a | 8 | 256 | 1.67 GMac | 1.60 M |
| H512_C2_a | 2 | 512 | 6.37 GMac | 6.34 M |
| H512_C4_a | 4 | 512 | 6.40 GMac | 6.34 M |
| H512_C8_a | 8 | 512 | 6.48 GMac | 6.35 M |
| Model | Channels | Hidden Dim | FLOPs | Parameters |
|---|---|---|---|---|
| H64_C2_b | 2 | 64 | 137.17 MMac | 119.01 k |
| H64_C4_b | 4 | 64 | 170.42 MMac | 119.14 k |
| H64_C8_b | 8 | 64 | 236.91 MMac | 119.39 k |
| H128_C2_b | 2 | 128 | 470.73 MMac | 434.59 k |
| H128_C4_b | 4 | 128 | 537.22 MMac | 434.85 k |
| H128_C8_b | 8 | 128 | 670.21 MMac | 435.36 k |
| H256_C2_b | 2 | 256 | 1.73 GMac | 1.66 M |
| H256_C4_b | 4 | 256 | 1.86 GMac | 1.66 M |
| H256_C8_b | 8 | 256 | 2.13 GMac | 1.66 M |
| H512_C2_b | 2 | 512 | 6.60 GMac | 6.46 M |
| H512_C4_b | 4 | 512 | 6.86 GMac | 6.46 M |
| H512_C8_b | 8 | 512 | 7.39 GMac | 6.46 M |
The implementation supports two main approaches:
- Input window size varies with latency (iW = LĂ—16)
- Output window size varies with latency (oW = LĂ—16)
- Fixed input window size (iW = 256 samples = 16ms)
- Output window size varies with latency (oW = LĂ—16)
- Low Latency: Supports latencies from 1ms to 16ms
- Lightweight: Models range from ~100k to ~25M parameters
- Multichannel: Supports 2, 4, and 8 channel configurations
- Time Domain: Operates directly on time-domain signals
- Efficient: Optimized for real-time processing on edge devices
Unlike traditional spatial attention, this model applies spatial_filters parameter.
The model uses learnable spatial filters to combine information across multiple input channels:
self.spatial_filters = nn.Parameter(torch.empty(hidden_dim, n_channels))
nn.init.kaiming_uniform_(self.spatial_filters)The model uses an overlap-add approach with configurable stride for efficient processing of long audio signals.
Layer normalization is applied before each LSTM layer to stabilize training and improve convergence.
If you use this code in your research, please cite the original paper:
@inproceedings{pandey2023simple,
title={A Simple RNN Model for Lightweight, Low-compute and Low-latency Multichannel Speech Enhancement in the Time Domain},
author={Pandey, Ashutosh and Tan, Ke and Xu, Buye},
booktitle={Proc. Interspeech 2023},
pages={2478--2482},
year={2023}
}This project is licensed under the terms specified in the LICENSE file.
This is an unofficial implementation of the paper. Please refer to the original paper for the official implementation and more details.
Contributions are welcome! Please feel free to submit a Pull Request.