Skip to content

About

🚀Unofficial implementation of "A Simple RNN Model for Lightweight, Low-compute and Low-latency Multichannel Speech Enhancement in the Time Domain" in PyTorch.

Resources

Stars

7 stars

Watchers

1 watching

Forks

Latest commit

 

History

9 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 

Repository files navigation

TimeDomain-RNN-SE

🚀 Unofficial implementation of A Simple RNN Model for Lightweight, Low-compute and Low-latency Multichannel Speech Enhancement in the Time Domain in PyTorch.

Overview

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.

Model Architecture

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

Key Parameters

  • 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)

Usage

Basic Usage

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)

Model Complexity Analysis

The repository includes comprehensive complexity analysis for various model configurations:

python simplernn.py

This 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 ($N=16000$).

Latency Explanation

Approach (a): Enhancement with Minimum Input Context

$iW = oW = L$

Approach (b): Enhancement with Fixed Input Context

$iW = 16\text{ms}$ (fixed), $oW = L$

Model Configurations

Table 3: Latency vs. Performance Trade-offs

Approach (a): Input and output windows vary with latency (iW = LĂ—16, oW = LĂ—16)

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

Table 2: Performance for 2ms Latency with Varying Widths

Performance of the proposed model for an algorithmic latency of 2 ms with varying widths.

Approach (a): Input window = 32 samples (2ms), Output window = 32 samples (2ms)

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

Approach (b): Input window = 256 samples (16ms), Output window = 32 samples (2ms)

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

Model Variants

The implementation supports two main approaches:

Approach (a)

  • Input window size varies with latency (iW = LĂ—16)
  • Output window size varies with latency (oW = LĂ—16)

Approach (b)

  • Fixed input window size (iW = 256 samples = 16ms)
  • Output window size varies with latency (oW = LĂ—16)

Features

  • 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

Technical Details

Spatial Filtering

Unlike traditional spatial attention, this model applies $H$ trainable linear filters across the $C$ channels to capture frequency-dependent spatial information efficiently: $$y_t = \sum_{c=1}^{C} w_{h,c} \cdot x_{t,c}$$ where $w \in \mathbb{R}^{H \times C}$ is the 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)

Overlap-Add Processing

The model uses an overlap-add approach with configurable stride for efficient processing of long audio signals.

Normalization

Layer normalization is applied before each LSTM layer to stabilize training and improve convergence.

Citation

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}
}

License

This project is licensed under the terms specified in the LICENSE file.

Acknowledgments

This is an unofficial implementation of the paper. Please refer to the original paper for the official implementation and more details.

Contributing

Contributions are welcome! Please feel free to submit a Pull Request.

About

🚀Unofficial implementation of "A Simple RNN Model for Lightweight, Low-compute and Low-latency Multichannel Speech Enhancement in the Time Domain" in PyTorch.

Resources

Stars

7 stars

Watchers

1 watching

Forks

Releases

Packages

Contributors

Languages