1use 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}