1use burn_core as burn;
2
3use crate::{
4 BatchNorm, BatchNormConfig, GroupNorm, GroupNormConfig, InstanceNorm, InstanceNormConfig,
5 LayerNorm, LayerNormConfig, RmsNorm, RmsNormConfig,
6};
7use burn::prelude::{Config, Module};
8use burn::tensor::Tensor;
9use burn::tensor::backend::Backend;
10
11#[derive(Config, Debug)]
20#[non_exhaustive]
21pub enum NormalizationConfig {
22 Batch(BatchNormConfig),
24
25 Group(GroupNormConfig),
27
28 Instance(InstanceNormConfig),
30
31 Layer(LayerNormConfig),
33
34 Rms(RmsNormConfig),
36}
37
38impl From<BatchNormConfig> for NormalizationConfig {
39 fn from(config: BatchNormConfig) -> Self {
40 Self::Batch(config)
41 }
42}
43
44impl From<GroupNormConfig> for NormalizationConfig {
45 fn from(config: GroupNormConfig) -> Self {
46 Self::Group(config)
47 }
48}
49
50impl From<InstanceNormConfig> for NormalizationConfig {
51 fn from(config: InstanceNormConfig) -> Self {
52 Self::Instance(config)
53 }
54}
55
56impl From<LayerNormConfig> for NormalizationConfig {
57 fn from(config: LayerNormConfig) -> Self {
58 Self::Layer(config)
59 }
60}
61
62impl From<RmsNormConfig> for NormalizationConfig {
63 fn from(config: RmsNormConfig) -> Self {
64 Self::Rms(config)
65 }
66}
67
68impl NormalizationConfig {
69 pub fn init<B: Backend>(&self, device: &B::Device) -> Normalization<B> {
71 match self {
72 NormalizationConfig::Batch(config) => config.init(device).into(),
73 NormalizationConfig::Group(config) => config.init(device).into(),
74 NormalizationConfig::Instance(config) => config.init(device).into(),
75 NormalizationConfig::Layer(config) => config.init(device).into(),
76 NormalizationConfig::Rms(config) => config.init(device).into(),
77 }
78 }
79
80 pub fn with_num_features(self, num_features: usize) -> Self {
82 match self {
83 NormalizationConfig::Batch(config) => BatchNormConfig {
84 num_features,
85 ..config
86 }
87 .into(),
88 NormalizationConfig::Group(config) => GroupNormConfig {
89 num_channels: num_features,
90 ..config
91 }
92 .into(),
93 NormalizationConfig::Instance(config) => InstanceNormConfig {
94 num_channels: num_features,
95 ..config
96 }
97 .into(),
98 NormalizationConfig::Layer(config) => LayerNormConfig {
99 d_model: num_features,
100 ..config
101 }
102 .into(),
103 NormalizationConfig::Rms(config) => RmsNormConfig {
104 d_model: num_features,
105 ..config
106 }
107 .into(),
108 }
109 }
110
111 pub fn num_features(&self) -> usize {
113 match self {
114 NormalizationConfig::Batch(config) => config.num_features,
115 NormalizationConfig::Group(config) => config.num_channels,
116 NormalizationConfig::Instance(config) => config.num_channels,
117 NormalizationConfig::Layer(config) => config.d_model,
118 NormalizationConfig::Rms(config) => config.d_model,
119 }
120 }
121}
122
123#[derive(Module, Debug)]
134#[non_exhaustive]
135pub enum Normalization<B: Backend> {
136 Batch(BatchNorm<B>),
138
139 Group(GroupNorm<B>),
141
142 Instance(InstanceNorm<B>),
144
145 Layer(LayerNorm<B>),
147
148 Rms(RmsNorm<B>),
150}
151
152impl<B: Backend> From<BatchNorm<B>> for Normalization<B> {
153 fn from(layer: BatchNorm<B>) -> Self {
154 Self::Batch(layer)
155 }
156}
157
158impl<B: Backend> From<GroupNorm<B>> for Normalization<B> {
159 fn from(layer: GroupNorm<B>) -> Self {
160 Self::Group(layer)
161 }
162}
163
164impl<B: Backend> From<InstanceNorm<B>> for Normalization<B> {
165 fn from(layer: InstanceNorm<B>) -> Self {
166 Self::Instance(layer)
167 }
168}
169
170impl<B: Backend> From<LayerNorm<B>> for Normalization<B> {
171 fn from(layer: LayerNorm<B>) -> Self {
172 Self::Layer(layer)
173 }
174}
175
176impl<B: Backend> From<RmsNorm<B>> for Normalization<B> {
177 fn from(layer: RmsNorm<B>) -> Self {
178 Self::Rms(layer)
179 }
180}
181
182impl<B: Backend> Normalization<B> {
183 pub fn forward<const D: usize>(&self, input: Tensor<B, D>) -> Tensor<B, D> {
189 match self {
190 Normalization::Batch(norm) => norm.forward(input),
191 Normalization::Group(norm) => norm.forward(input),
192 Normalization::Instance(norm) => norm.forward(input),
193 Normalization::Layer(norm) => norm.forward(input),
194 Normalization::Rms(norm) => norm.forward(input),
195 }
196 }
197
198 pub fn num_features(&self) -> usize {
200 match self {
201 Normalization::Batch(norm) => norm.gamma.shape().dims[0],
202 Normalization::Group(norm) => norm.num_channels,
203 Normalization::Instance(norm) => norm.num_channels,
204 Normalization::Layer(norm) => norm.gamma.shape().dims[0],
205 Normalization::Rms(norm) => norm.gamma.shape().dims[0],
206 }
207 }
208}
209
210#[cfg(feature = "std")]
211#[cfg(test)]
212mod tests {
213 use super::*;
214 use crate::TestAutodiffBackend;
215
216 #[test]
217 fn test_match_feature_size() {
218 let config: NormalizationConfig = BatchNormConfig::new(0).into();
219 assert_eq!(config.num_features(), 0);
220 let config = config.with_num_features(12);
221 assert_eq!(config.num_features(), 12);
222
223 let config: NormalizationConfig = GroupNormConfig::new(4, 0).into();
224 assert_eq!(config.num_features(), 0);
225 let config = config.with_num_features(12);
226 assert_eq!(config.num_features(), 12);
227
228 let config: NormalizationConfig = InstanceNormConfig::new(0).into();
229 assert_eq!(config.num_features(), 0);
230 let config = config.with_num_features(12);
231 assert_eq!(config.num_features(), 12);
232
233 let config: NormalizationConfig = LayerNormConfig::new(0).into();
234 assert_eq!(config.num_features(), 0);
235 let config = config.with_num_features(12);
236 assert_eq!(config.num_features(), 12);
237
238 let config: NormalizationConfig = RmsNormConfig::new(0).into();
239 assert_eq!(config.num_features(), 0);
240 let config = config.with_num_features(12);
241 assert_eq!(config.num_features(), 12);
242 }
243
244 #[test]
245 fn test_batch_norm() {
246 type B = TestAutodiffBackend;
247 let device = Default::default();
248
249 let num_features = 12;
250 let input: Tensor<B, 4> = Tensor::ones([2, num_features, 3, 4], &device);
251
252 let config: NormalizationConfig = BatchNormConfig::new(12).into();
253
254 let layer: Normalization<B> = config.init(&device);
255 assert_eq!(layer.num_features(), 12);
256
257 let expected = match &layer {
258 Normalization::Batch(inner) => inner.forward(input.clone()),
259 _ => panic!("Unexpected layer type"),
260 };
261
262 let output = layer.forward(input);
263
264 output.to_data().assert_eq(&expected.to_data(), true);
265 }
266
267 #[test]
268 fn test_group_norm() {
269 type B = TestAutodiffBackend;
270 let device = Default::default();
271
272 let num_features = 12;
273 let input: Tensor<B, 4> = Tensor::ones([2, num_features, 3, 4], &device);
274
275 let config: NormalizationConfig = GroupNormConfig::new(3, num_features).into();
276
277 let layer: Normalization<B> = config.init(&device);
278 assert_eq!(layer.num_features(), 12);
279
280 let expected = match &layer {
281 Normalization::Group(inner) => inner.forward(input.clone()),
282 _ => panic!("Unexpected layer type"),
283 };
284
285 let output = layer.forward(input);
286
287 output.to_data().assert_eq(&expected.to_data(), true);
288 }
289
290 #[test]
291 fn test_instance_norm() {
292 type B = TestAutodiffBackend;
293 let device = Default::default();
294
295 let num_features = 12;
296 let input: Tensor<B, 4> = Tensor::ones([2, num_features, 3, 4], &device);
297
298 let config: NormalizationConfig = InstanceNormConfig::new(num_features).into();
299
300 let layer: Normalization<B> = config.init(&device);
301 assert_eq!(layer.num_features(), 12);
302
303 let expected = match &layer {
304 Normalization::Instance(inner) => inner.forward(input.clone()),
305 _ => panic!("Unexpected layer type"),
306 };
307
308 let output = layer.forward(input);
309
310 output.to_data().assert_eq(&expected.to_data(), true);
311 }
312
313 #[test]
314 fn test_layer_norm() {
315 type B = TestAutodiffBackend;
316 let device = Default::default();
317
318 let num_features = 12;
319 let input: Tensor<B, 4> = Tensor::ones([2, 3, 4, num_features], &device);
320
321 let config: NormalizationConfig = LayerNormConfig::new(num_features).into();
322
323 let layer: Normalization<B> = config.init(&device);
324 assert_eq!(layer.num_features(), 12);
325
326 let expected = match &layer {
327 Normalization::Layer(inner) => inner.forward(input.clone()),
328 _ => panic!("Unexpected layer type"),
329 };
330
331 let output = layer.forward(input);
332
333 output.to_data().assert_eq(&expected.to_data(), true);
334 }
335
336 #[test]
337 fn test_rms_norm() {
338 type B = TestAutodiffBackend;
339 let device = Default::default();
340
341 let num_features = 12;
342 let input: Tensor<B, 4> = Tensor::ones([2, 3, 4, num_features], &device);
343
344 let config: NormalizationConfig = RmsNormConfig::new(num_features).into();
345
346 let layer: Normalization<B> = config.init(&device);
347 assert_eq!(layer.num_features(), 12);
348
349 let expected = match &layer {
350 Normalization::Rms(inner) => inner.forward(input.clone()),
351 _ => panic!("Unexpected layer type"),
352 };
353
354 let output = layer.forward(input);
355
356 output.to_data().assert_eq(&expected.to_data(), true);
357 }
358}