voxtral_micro/tts/codec/
layer_scale.rs1use burn::module::{Param, ParamId};
7use burn::tensor::backend::Backend;
8use burn::tensor::Tensor;
9
10#[derive(burn::module::Module, Debug)]
19pub struct LayerScale<B: Backend> {
20 pub scale: Param<Tensor<B, 1>>,
22}
23
24impl<B: Backend> LayerScale<B> {
25 pub fn new(scale: Tensor<B, 1>) -> Self {
30 Self {
31 scale: Param::initialized(ParamId::new(), scale),
32 }
33 }
34
35 pub fn forward(&self, x: Tensor<B, 3>) -> Tensor<B, 3> {
43 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 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 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 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 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 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 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 let x = Tensor::<TestBackend, 3>::from_data(
164 TensorData::new(
165 vec![
166 1.0f32, 1.0, 1.0, 1.0, 3.0, 3.0, 3.0, 3.0, ],
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 assert!((vals[0] - 2.0).abs() < 1e-6);
180 assert!((vals[4] - 6.0).abs() < 1e-6);
182 }
183}