Skip to main content

voxtral_micro/tts/codec/
layer_scale.rs

1//! LayerScale: learnable per-channel scaling for residual connections.
2//!
3//! Used in the codec decoder to scale attention and FFN outputs before
4//! adding to the residual stream.
5
6use burn::module::{Param, ParamId};
7use burn::tensor::backend::Backend;
8use burn::tensor::Tensor;
9
10/// LayerScale applies a learnable per-channel scale to its input.
11///
12/// Forward: `output = input * scale` where scale is [dim] broadcast over
13/// batch and sequence dimensions.
14///
15/// In the codec decoder, each transformer layer has two LayerScale instances:
16/// - `attention_scale` [1024]: scales attention output before residual add
17/// - `ffn_scale` [1024]: scales FFN output before residual add
18#[derive(burn::module::Module, Debug)]
19pub struct LayerScale<B: Backend> {
20    /// Per-channel scale weights [dim].
21    pub scale: Param<Tensor<B, 1>>,
22}
23
24impl<B: Backend> LayerScale<B> {
25    /// Create LayerScale from a loaded weight tensor.
26    ///
27    /// # Arguments
28    /// * `scale` - Scale weights [dim], typically initialized to 0.01 during training.
29    pub fn new(scale: Tensor<B, 1>) -> Self {
30        Self {
31            scale: Param::initialized(ParamId::new(), scale),
32        }
33    }
34
35    /// Apply per-channel scaling to input.
36    ///
37    /// # Arguments
38    /// * `x` - Input tensor [batch, seq_len, dim]
39    ///
40    /// # Returns
41    /// Scaled tensor [batch, seq_len, dim].
42    pub fn forward(&self, x: Tensor<B, 3>) -> Tensor<B, 3> {
43        // scale [dim] broadcasts over [batch, seq_len, dim]
44        x * self.scale.val().unsqueeze::<3>().unsqueeze()
45    }
46}
47
48#[cfg(test)]
49mod tests {
50    use super::*;
51    use burn::backend::Wgpu;
52    use burn::tensor::TensorData;
53
54    type TestBackend = Wgpu;
55
56    #[test]
57    fn test_layer_scale_output_shape() {
58        let device = Default::default();
59        let dim = 1024;
60
61        let scale = Tensor::<TestBackend, 1>::ones([dim], &device) * 0.01;
62        let ls = LayerScale::new(scale);
63
64        let x = Tensor::<TestBackend, 3>::ones([2, 10, dim], &device);
65        let out = ls.forward(x);
66
67        assert_eq!(out.dims(), [2, 10, dim]);
68    }
69
70    #[test]
71    fn test_layer_scale_multiplies_correctly() {
72        let device = Default::default();
73        let dim = 4;
74
75        // Scale = [1, 2, 3, 4]
76        let scale = Tensor::<TestBackend, 1>::from_data(
77            TensorData::new(vec![1.0f32, 2.0, 3.0, 4.0], [4]),
78            &device,
79        );
80        let ls = LayerScale::new(scale);
81
82        // Input = ones [1, 2, 4]
83        let x = Tensor::<TestBackend, 3>::ones([1, 2, dim], &device);
84        let out = ls.forward(x);
85
86        let data = out.to_data();
87        let vals = data.as_slice::<f32>().unwrap();
88
89        // Each position should be scaled by [1, 2, 3, 4]
90        // Position (0, 0, :) = [1, 2, 3, 4]
91        assert!((vals[0] - 1.0).abs() < 1e-6);
92        assert!((vals[1] - 2.0).abs() < 1e-6);
93        assert!((vals[2] - 3.0).abs() < 1e-6);
94        assert!((vals[3] - 4.0).abs() < 1e-6);
95
96        // Position (0, 1, :) should be the same (broadcast)
97        assert!((vals[4] - 1.0).abs() < 1e-6);
98        assert!((vals[5] - 2.0).abs() < 1e-6);
99        assert!((vals[6] - 3.0).abs() < 1e-6);
100        assert!((vals[7] - 4.0).abs() < 1e-6);
101    }
102
103    #[test]
104    fn test_layer_scale_default_init_value() {
105        let device = Default::default();
106        let dim = 1024;
107
108        // Codec layers are initialized to 0.01 during training
109        let scale = Tensor::<TestBackend, 1>::ones([dim], &device) * 0.01;
110        let ls = LayerScale::new(scale);
111
112        let x = Tensor::<TestBackend, 3>::ones([1, 5, dim], &device);
113        let out = ls.forward(x);
114
115        let data = out.to_data();
116        let vals = data.as_slice::<f32>().unwrap();
117
118        // All values should be ~0.01
119        for (i, &v) in vals.iter().enumerate().take(10) {
120            assert!(
121                (v - 0.01).abs() < 1e-6,
122                "Value[{}] = {} (expected 0.01)",
123                i,
124                v
125            );
126        }
127    }
128
129    #[test]
130    fn test_layer_scale_zero_scale_zeros_output() {
131        let device = Default::default();
132        let dim = 8;
133
134        let scale = Tensor::<TestBackend, 1>::zeros([dim], &device);
135        let ls = LayerScale::new(scale);
136
137        let x = Tensor::<TestBackend, 3>::ones([1, 3, dim], &device) * 42.0;
138        let out = ls.forward(x);
139
140        let data = out.to_data();
141        let vals = data.as_slice::<f32>().unwrap();
142
143        for (i, &v) in vals.iter().enumerate() {
144            assert!(
145                v.abs() < 1e-7,
146                "Zero scale should zero output, got val[{}] = {}",
147                i,
148                v
149            );
150        }
151    }
152
153    #[test]
154    fn test_layer_scale_batch_independence() {
155        let device = Default::default();
156        let scale = Tensor::<TestBackend, 1>::from_data(
157            TensorData::new(vec![2.0f32, 2.0, 2.0, 2.0], [4]),
158            &device,
159        );
160        let ls = LayerScale::new(scale);
161
162        // Different batch items should be scaled identically
163        let x = Tensor::<TestBackend, 3>::from_data(
164            TensorData::new(
165                vec![
166                    1.0f32, 1.0, 1.0, 1.0, // batch 0, seq 0
167                    3.0, 3.0, 3.0, 3.0, // batch 1, seq 0
168                ],
169                [2, 1, 4],
170            ),
171            &device,
172        );
173        let out = ls.forward(x);
174
175        let data = out.to_data();
176        let vals = data.as_slice::<f32>().unwrap();
177
178        // batch 0: [1,1,1,1] * 2 = [2,2,2,2]
179        assert!((vals[0] - 2.0).abs() < 1e-6);
180        // batch 1: [3,3,3,3] * 2 = [6,6,6,6]
181        assert!((vals[4] - 6.0).abs() < 1e-6);
182    }
183}