1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
//! # burn-gdn2
//!
//! Gated DeltaNet 2 (GDN-2) — a linear‑complexity recurrent token mixer
//! with channel‑wise erase/write gates.
//!
//! ## Quick start
//!
//! ```rust
//! use burn::backend::NdArray;
//! use burn::tensor::Tensor;
//! use burn_gdn2::{Gdn2Config, Gdn2Mode, GatedDeltaNet2};
//!
//! let device = burn::backend::ndarray::NdArrayDevice::Cpu;
//! let config = Gdn2Config {
//! hidden_size: 64,
//! num_heads: 2,
//! head_dim: 32,
//! ..Default::default()
//! };
//! let model = GatedDeltaNet2::<NdArray>::new(&config, &device);
//!
//! // Inference — token‑by‑token, state passed by reference
//! let input = Tensor::zeros([1, 16, 64], &device);
//! let mut state = None;
//! let output = model.forward(input, &mut state, true);
//!
//! // Training — full sequence, chunked WY for efficiency
//! let input = Tensor::zeros([1, 128, 64], &device);
//! let output = model.forward_train(input);
//! ```
//!
//! ## Features
//!
//! - **`std`** (default) — standard library support
//! - **`autodiff`** — differentiation support (required for training)
//! - **`cubecl`** — CubeCL‑accelerated GPU kernels (experimental)
pub use ;
pub use ;
pub use l2_normalize;
pub use ;
pub use short_conv_1d;