Skip to main content

nam_rs/models/
nam_model.rs

1// SPDX-License-Identifier: Apache-2.0
2// Copyright (c) 2026 Fábio Henrique de Lima Silva (fhl.bsb@gmail.com) All rights reserved.
3
4use super::slimmable::SlimmableModel;
5use super::{NamModel, StaticModel};
6
7impl NamModel for StaticModel {
8    #[inline(always)]
9    fn process(&mut self, input: &[f32], output: &mut [f32]) {
10        match self {
11            Self::WavenetStandard(m) => m.process(input, output),
12            Self::WavenetLite(m) => m.process(input, output),
13            Self::WavenetFeather(m) => m.process(input, output),
14            Self::WavenetNano(m) => m.process(input, output),
15            Self::WavenetA2Full(m) => m.process(input, output),
16            Self::WavenetA2Lite(m) => m.process(input, output),
17            Self::WavenetA2Dyn(m) => m.process(input, output),
18            Self::WavenetA2Cascade(m) => m.process(input, output),
19            Self::WavenetDyn(m) => m.process(input, output),
20            Self::Container(m) => m.process(input, output),
21            Self::Lstm1x3(m) => m.process(input, output),
22            Self::Lstm1x8(m) => m.process(input, output),
23            Self::Lstm1x12(m) => m.process(input, output),
24            Self::Lstm1x16(m) => m.process(input, output),
25            Self::Lstm1x24(m) => m.process(input, output),
26            Self::Lstm2x8(m) => m.process(input, output),
27            Self::Lstm2x12(m) => m.process(input, output),
28            Self::Lstm2x16(m) => m.process(input, output),
29            Self::Lstm1x40(m) => m.process(input, output),
30            Self::Lstm2x24(m) => m.process(input, output),
31            Self::LstmDyn(m) => m.process(input, output),
32            Self::Linear(m) => unsafe { m.process(input, output) },
33            Self::ConvNet(m) => m.process(input, output),
34        }
35    }
36
37    #[cold]
38    fn prewarm(&mut self, num_samples: usize) {
39        match self {
40            Self::WavenetStandard(m) => m.prewarm(),
41            Self::WavenetLite(m) => m.prewarm(),
42            Self::WavenetFeather(m) => m.prewarm(),
43            Self::WavenetNano(m) => m.prewarm(),
44            Self::WavenetA2Full(m) => m.prewarm(),
45            Self::WavenetA2Lite(m) => m.prewarm(),
46            Self::WavenetA2Dyn(m) => m.prewarm(),
47            Self::WavenetA2Cascade(m) => m.prewarm(),
48            Self::WavenetDyn(m) => m.prewarm(),
49            Self::Container(m) => m.prewarm(num_samples),
50            Self::Lstm1x3(m) => m.prewarm(num_samples),
51            Self::Lstm1x8(m) => m.prewarm(num_samples),
52            Self::Lstm1x12(m) => m.prewarm(num_samples),
53            Self::Lstm1x16(m) => m.prewarm(num_samples),
54            Self::Lstm1x24(m) => m.prewarm(num_samples),
55            Self::Lstm2x8(m) => m.prewarm(num_samples),
56            Self::Lstm2x12(m) => m.prewarm(num_samples),
57            Self::Lstm2x16(m) => m.prewarm(num_samples),
58            Self::Lstm1x40(m) => m.prewarm(num_samples),
59            Self::Lstm2x24(m) => m.prewarm(num_samples),
60            Self::LstmDyn(m) => m.prewarm(num_samples),
61            Self::Linear(m) => m.prewarm(num_samples),
62            Self::ConvNet(m) => m.prewarm(),
63        }
64    }
65
66    fn prewarm_on_reset(&self) -> bool {
67        match self {
68            Self::WavenetStandard(m) => m.prewarm_on_reset(),
69            Self::WavenetLite(m) => m.prewarm_on_reset(),
70            Self::WavenetFeather(m) => m.prewarm_on_reset(),
71            Self::WavenetNano(m) => m.prewarm_on_reset(),
72            Self::WavenetA2Full(m) => m.prewarm_on_reset(),
73            Self::WavenetA2Lite(m) => m.prewarm_on_reset(),
74            Self::WavenetA2Dyn(m) => m.prewarm_on_reset(),
75            Self::WavenetA2Cascade(m) => m.prewarm_on_reset(),
76            Self::WavenetDyn(m) => m.prewarm_on_reset(),
77            Self::Container(m) => m.prewarm_on_reset(),
78            Self::Lstm1x3(m) => m.prewarm_on_reset(),
79            Self::Lstm1x8(m) => m.prewarm_on_reset(),
80            Self::Lstm1x12(m) => m.prewarm_on_reset(),
81            Self::Lstm1x16(m) => m.prewarm_on_reset(),
82            Self::Lstm1x24(m) => m.prewarm_on_reset(),
83            Self::Lstm2x8(m) => m.prewarm_on_reset(),
84            Self::Lstm2x12(m) => m.prewarm_on_reset(),
85            Self::Lstm2x16(m) => m.prewarm_on_reset(),
86            Self::Lstm1x40(m) => m.prewarm_on_reset(),
87            Self::Lstm2x24(m) => m.prewarm_on_reset(),
88            Self::LstmDyn(m) => m.prewarm_on_reset(),
89            Self::Linear(m) => m.prewarm_on_reset(),
90            Self::ConvNet(m) => m.prewarm_on_reset(),
91        }
92    }
93
94    fn set_prewarm_on_reset(&mut self, val: bool) {
95        match self {
96            Self::WavenetStandard(m) => m.set_prewarm_on_reset(val),
97            Self::WavenetLite(m) => m.set_prewarm_on_reset(val),
98            Self::WavenetFeather(m) => m.set_prewarm_on_reset(val),
99            Self::WavenetNano(m) => m.set_prewarm_on_reset(val),
100            Self::WavenetA2Full(m) => m.set_prewarm_on_reset(val),
101            Self::WavenetA2Lite(m) => m.set_prewarm_on_reset(val),
102            Self::WavenetA2Dyn(m) => m.set_prewarm_on_reset(val),
103            Self::WavenetA2Cascade(m) => m.set_prewarm_on_reset(val),
104            Self::WavenetDyn(m) => m.set_prewarm_on_reset(val),
105            Self::Container(m) => m.set_prewarm_on_reset(val),
106            Self::Lstm1x3(m) => m.set_prewarm_on_reset(val),
107            Self::Lstm1x8(m) => m.set_prewarm_on_reset(val),
108            Self::Lstm1x12(m) => m.set_prewarm_on_reset(val),
109            Self::Lstm1x16(m) => m.set_prewarm_on_reset(val),
110            Self::Lstm1x24(m) => m.set_prewarm_on_reset(val),
111            Self::Lstm2x8(m) => m.set_prewarm_on_reset(val),
112            Self::Lstm2x12(m) => m.set_prewarm_on_reset(val),
113            Self::Lstm2x16(m) => m.set_prewarm_on_reset(val),
114            Self::Lstm1x40(m) => m.set_prewarm_on_reset(val),
115            Self::Lstm2x24(m) => m.set_prewarm_on_reset(val),
116            Self::LstmDyn(m) => m.set_prewarm_on_reset(val),
117            Self::Linear(m) => m.set_prewarm_on_reset(val),
118            Self::ConvNet(m) => m.set_prewarm_on_reset(val),
119        }
120    }
121
122    fn reset(&mut self, sample_rate: u32, max_buffer_size: usize) -> anyhow::Result<()> {
123        match self {
124            Self::WavenetStandard(m) => m.reset(sample_rate, max_buffer_size),
125            Self::WavenetLite(m) => m.reset(sample_rate, max_buffer_size),
126            Self::WavenetFeather(m) => m.reset(sample_rate, max_buffer_size),
127            Self::WavenetNano(m) => m.reset(sample_rate, max_buffer_size),
128            Self::WavenetA2Full(m) => m.reset(sample_rate, max_buffer_size),
129            Self::WavenetA2Lite(m) => m.reset(sample_rate, max_buffer_size),
130            Self::WavenetA2Dyn(m) => m.reset(sample_rate, max_buffer_size),
131            Self::WavenetA2Cascade(m) => m.reset(sample_rate, max_buffer_size),
132            Self::WavenetDyn(m) => m.reset(sample_rate, max_buffer_size),
133            Self::Container(m) => m.reset(sample_rate, max_buffer_size),
134            Self::Lstm1x3(m) => m.reset(sample_rate, max_buffer_size),
135            Self::Lstm1x8(m) => m.reset(sample_rate, max_buffer_size),
136            Self::Lstm1x12(m) => m.reset(sample_rate, max_buffer_size),
137            Self::Lstm1x16(m) => m.reset(sample_rate, max_buffer_size),
138            Self::Lstm1x24(m) => m.reset(sample_rate, max_buffer_size),
139            Self::Lstm2x8(m) => m.reset(sample_rate, max_buffer_size),
140            Self::Lstm2x12(m) => m.reset(sample_rate, max_buffer_size),
141            Self::Lstm2x16(m) => m.reset(sample_rate, max_buffer_size),
142            Self::Lstm1x40(m) => m.reset(sample_rate, max_buffer_size),
143            Self::Lstm2x24(m) => m.reset(sample_rate, max_buffer_size),
144            Self::LstmDyn(m) => m.reset(sample_rate, max_buffer_size),
145            Self::Linear(m) => NamModel::reset(m.as_mut(), sample_rate, max_buffer_size),
146            Self::ConvNet(m) => NamModel::reset(m.as_mut(), sample_rate, max_buffer_size),
147        }
148    }
149
150    fn set_max_buffer_size(&mut self, max_buf: usize) -> anyhow::Result<()> {
151        match self {
152            Self::WavenetStandard(m) => m.set_max_buffer_size(max_buf),
153            Self::WavenetLite(m) => m.set_max_buffer_size(max_buf),
154            Self::WavenetFeather(m) => m.set_max_buffer_size(max_buf),
155            Self::WavenetNano(m) => m.set_max_buffer_size(max_buf),
156            Self::WavenetA2Full(m) => m.set_max_buffer_size(max_buf),
157            Self::WavenetA2Lite(m) => m.set_max_buffer_size(max_buf),
158            Self::WavenetA2Dyn(m) => m.set_max_buffer_size(max_buf),
159            Self::WavenetA2Cascade(m) => m.set_max_buffer_size(max_buf),
160            Self::WavenetDyn(m) => m.set_max_buffer_size(max_buf),
161            Self::Container(m) => m.set_max_buffer_size(max_buf),
162            Self::Lstm1x3(m) => m.set_max_buffer_size(max_buf),
163            Self::Lstm1x8(m) => m.set_max_buffer_size(max_buf),
164            Self::Lstm1x12(m) => m.set_max_buffer_size(max_buf),
165            Self::Lstm1x16(m) => m.set_max_buffer_size(max_buf),
166            Self::Lstm1x24(m) => m.set_max_buffer_size(max_buf),
167            Self::Lstm2x8(m) => m.set_max_buffer_size(max_buf),
168            Self::Lstm2x12(m) => m.set_max_buffer_size(max_buf),
169            Self::Lstm2x16(m) => m.set_max_buffer_size(max_buf),
170            Self::Lstm1x40(m) => m.set_max_buffer_size(max_buf),
171            Self::Lstm2x24(m) => m.set_max_buffer_size(max_buf),
172            Self::LstmDyn(m) => m.set_max_buffer_size(max_buf),
173            Self::Linear(m) => NamModel::set_max_buffer_size(m.as_mut(), max_buf),
174            Self::ConvNet(m) => NamModel::set_max_buffer_size(m.as_mut(), max_buf),
175        }
176    }
177
178    fn prewarm_samples(&self) -> usize {
179        match self {
180            Self::WavenetStandard(m) => m.prewarm_samples(),
181            Self::WavenetLite(m) => m.prewarm_samples(),
182            Self::WavenetFeather(m) => m.prewarm_samples(),
183            Self::WavenetNano(m) => m.prewarm_samples(),
184            Self::WavenetA2Full(m) => m.prewarm_samples(),
185            Self::WavenetA2Lite(m) => m.prewarm_samples(),
186            Self::WavenetA2Dyn(m) => m.prewarm_samples(),
187            Self::WavenetA2Cascade(m) => m.prewarm_samples(),
188            Self::WavenetDyn(m) => m.prewarm_samples(),
189            Self::Container(m) => m.prewarm_samples(),
190            Self::Lstm1x3(m) => m.prewarm_samples(),
191            Self::Lstm1x8(m) => m.prewarm_samples(),
192            Self::Lstm1x12(m) => m.prewarm_samples(),
193            Self::Lstm1x16(m) => m.prewarm_samples(),
194            Self::Lstm1x24(m) => m.prewarm_samples(),
195            Self::Lstm2x8(m) => m.prewarm_samples(),
196            Self::Lstm2x12(m) => m.prewarm_samples(),
197            Self::Lstm2x16(m) => m.prewarm_samples(),
198            Self::Lstm1x40(m) => m.prewarm_samples(),
199            Self::Lstm2x24(m) => m.prewarm_samples(),
200            Self::LstmDyn(m) => m.prewarm_samples(),
201            Self::Linear(m) => m.prewarm_samples(),
202            Self::ConvNet(m) => m.prewarm_samples(),
203        }
204    }
205
206    fn slimmable_breakpoints(&self) -> Vec<f64> {
207        if let Self::Container(c) = self {
208            SlimmableModel::slimmable_breakpoints(c.as_ref())
209        } else {
210            vec![]
211        }
212    }
213}