1pub mod conv;
2pub mod loader;
3
4use anyhow::{Context, Result};
5use burn::module::{Module, Param};
6use burn::nn::PaddingConfig1d;
7use burn::nn::conv::{Conv1d, Conv1dConfig};
8use burn::nn::{LayerNorm, LayerNormConfig, Linear, LinearConfig, LinearLayout};
9use burn::prelude::Backend;
10use burn::tensor::{DType, Int, Tensor, TensorData};
11use burn_store::{BurnpackStore, ModuleSnapshot, ModuleStore};
12use oxideav_core::{
13 CodecId, CodecParameters, MediaType, Packet, RuntimeContext, SampleFormat, StreamInfo, TimeBase,
14};
15use rayon::prelude::*;
16
17pub use conv::PostProcessor;
18pub use conv::{PlainConv1d, WNConv1d, WNConvTranspose1d};
19
20const HEARTMULA_SAMPLE_RATE: usize = 48_000;
21const HEARTCODEC_WINDOW_FRAMES: usize = 93;
22const HEARTCODEC_SEGMENT_DURATION_SECONDS: f32 = 29.76;
23type TensorLookup = dyn Fn(&str) -> Option<(Vec<f32>, Vec<usize>)>;
24
25#[derive(Debug, Clone)]
26pub struct HeartCodecConfig {
27 pub dim: usize,
28 pub codebook_size: usize,
29 pub codebook_dim: usize,
30 pub num_quantizers: usize,
31 pub attention_head_dim: usize,
32 pub in_channels: usize,
33 pub num_attention_heads: usize,
34 pub num_layers: usize,
35 pub num_layers_2: usize,
36 pub out_channels: usize,
37 pub sample_rate: usize,
38 pub latent_hidden_dim: usize,
39 pub init_channel: usize,
40 pub num_bands: usize,
41 pub num_samples: usize,
42 pub downsample_factors: [usize; 5],
43 pub downsample_kernel_sizes: [usize; 5],
44 pub upsample_factors: [usize; 5],
45 pub upsample_kernel_sizes: [usize; 5],
46 pub default_kernel_size: usize,
47 pub delay_kernel_size: usize,
48 pub res_kernel_size: usize,
49 pub causal: bool,
50
51 pub ode_steps: usize,
52}
53
54impl Default for HeartCodecConfig {
55 fn default() -> Self {
56 Self {
57 dim: 512,
58 codebook_size: 8192,
59 codebook_dim: 32,
60 num_quantizers: 8,
61 attention_head_dim: 64,
62 in_channels: 1024,
63 num_attention_heads: 24,
64 num_layers: 24,
65 num_layers_2: 6,
66 out_channels: 256,
67 sample_rate: 48000,
68 latent_hidden_dim: 128,
69 init_channel: 64,
70 num_bands: 1,
71 num_samples: 2,
72 downsample_factors: [3, 4, 4, 4, 5],
73 downsample_kernel_sizes: [6, 8, 8, 8, 10],
74 upsample_factors: [5, 4, 4, 4, 3],
75 upsample_kernel_sizes: [10, 8, 8, 8, 6],
76 default_kernel_size: 7,
77 delay_kernel_size: 5,
78 res_kernel_size: 7,
79 causal: true,
80 ode_steps: 10,
81 }
82 }
83}
84
85#[derive(Module, Debug)]
86pub struct HeartCodecModel<B: Backend> {
87 pub flow_matching: FlowMatching<B>,
88 pub scalar_model: ScalarModel<B>,
89 pub ode_steps: usize,
90 pub guidance_scale: f32,
91}
92
93#[derive(Debug)]
94pub struct ScalarDecodePlan<B: Backend> {
95 pub target_len: usize,
96 pub audio_target_len: usize,
97 pub windows: Vec<Tensor<B, 3>>,
98}
99
100impl<B: Backend> HeartCodecModel<B> {
101 pub fn new(device: &B::Device) -> Self {
102 let config = HeartCodecConfig::default();
103 Self {
104 flow_matching: FlowMatching::new(device, &config),
105 scalar_model: ScalarModel::new(device, &config),
106 ode_steps: config.ode_steps,
107 guidance_scale: 1.0,
108 }
109 }
110
111 pub fn with_ode_steps(mut self, steps: usize) -> Self {
112 self.ode_steps = steps.clamp(1, 50);
113 self
114 }
115
116 pub fn with_guidance_scale(mut self, scale: f32) -> Self {
117 self.guidance_scale = scale.max(1.0);
118 self
119 }
120
121 pub fn from_burnpack(path: &std::path::Path, device: &B::Device) -> Result<Self> {
122 let mut model = Self::new(device);
123 let mut store = BurnpackStore::from_file(path).zero_copy(true);
124
125 if let Err(_e) = model.load_from(&mut store) {
126 model = Self::load_with_mapping(path, device)?;
127 }
128
129 Ok(model)
130 }
131
132 fn load_flow_matching_manually<F>(
133 _flow_matching: &mut FlowMatching<B>,
134 _path: &std::path::Path,
135 device: &B::Device,
136 get_tensor: &F,
137 ) -> Result<FlowMatching<B>>
138 where
139 F: Fn(&str) -> Option<(Vec<f32>, Vec<usize>)>,
140 {
141 let config = HeartCodecConfig::default();
142
143 let cond_feature_emb = load_linear_from_tensors(
144 device,
145 get_tensor,
146 "flow_matching.cond_feature_emb",
147 config.dim,
148 config.dim,
149 )?;
150 let zero_cond_embedding1 =
151 if let Some((data, shape)) = get_tensor("flow_matching.zero_cond_embedding1") {
152 Tensor::<B, 1>::from_data(TensorData::new(data, [shape[0]]), device)
153 } else {
154 Tensor::zeros([config.dim], device)
155 };
156
157 let vq_embed = ResidualVQ::load_from_dot_notation(device, get_tensor)?;
158
159 let estimator = LlamaTransformer::load_from_burnpack(device, get_tensor)?;
160
161 Ok(FlowMatching {
162 cond_feature_emb,
163 zero_cond_embedding1: Param::from_tensor(zero_cond_embedding1),
164 estimator,
165 vq_embed,
166 debug_latent_steps: false,
167 })
168 }
169
170 fn load_with_mapping(path: &std::path::Path, device: &B::Device) -> Result<Self> {
171 use burn::tensor::DType;
172 use burn_store::ModuleStore;
173
174 let snapshots = BurnpackStore::from_file(path)
175 .zero_copy(true)
176 .get_all_snapshots()
177 .with_context(|| "failed to read burnpack snapshots")?
178 .clone();
179
180 let mut model = Self::new(device);
181
182 let get_tensor = |name: &str| -> Option<(Vec<f32>, Vec<usize>)> {
183 snapshots.iter().find_map(|(_, snap)| {
184 if snap.full_path() == name {
185 let data = snap.to_data().ok()?;
186 if data.dtype == DType::F32 {
187 let shape = data.shape.to_vec();
188 let values: Vec<f32> = data.to_vec::<f32>().ok()?;
189 Some((values, shape))
190 } else {
191 None
192 }
193 } else {
194 None
195 }
196 })
197 };
198
199 match Self::load_flow_matching_manually(&mut model.flow_matching, path, device, &get_tensor)
200 {
201 Ok(flow_matching) => {
202 model.flow_matching = flow_matching;
203 }
204 Err(_e) => {}
205 }
206
207 match ScalarModel::load_from_dot_notation(path, device, &get_tensor) {
208 Ok(scalar_model) => {
209 model.scalar_model = scalar_model;
210 }
211 Err(_e) => {}
212 }
213
214 Ok(model)
215 }
216
217 fn snapshots_to_f32_lookup(path: &std::path::Path) -> Result<Box<TensorLookup>> {
218 let snapshots = BurnpackStore::from_file(path)
219 .zero_copy(true)
220 .get_all_snapshots()
221 .with_context(|| "failed to read burnpack snapshots")?
222 .clone();
223
224 Ok(Box::new(
225 move |name: &str| -> Option<(Vec<f32>, Vec<usize>)> {
226 snapshots.iter().find_map(|(_, snap)| {
227 if snap.full_path() == name {
228 let data = snap.to_data().ok()?;
229 if data.dtype == DType::F32 {
230 let shape = data.shape.to_vec();
231 let values: Vec<f32> = data.to_vec::<f32>().ok()?;
232 Some((values, shape))
233 } else {
234 None
235 }
236 } else {
237 None
238 }
239 })
240 },
241 ))
242 }
243
244 pub(crate) fn build_scalar_decode_plan_impl(
245 flow_matching: &FlowMatching<B>,
246 guidance_scale: f32,
247 ode_steps: usize,
248 codes: Tensor<B, 3, Int>,
249 first_latent: Tensor<B, 3>,
250 ) -> ScalarDecodePlan<B> {
251 let device = codes.device();
252 let [batch, num_quantizers, seq_len] = codes.dims();
253 assert_eq!(batch, 1, "HeartCodec decode expects a single code batch");
254
255 let duration_seconds = seq_len as f32 / 12.5;
256 let segment_duration_seconds = HEARTCODEC_SEGMENT_DURATION_SECONDS;
257 let latent_length = (segment_duration_seconds * 25.0) as usize;
258 assert_eq!(
259 first_latent.dims(),
260 [batch, latent_length, 256],
261 "initial latent shape must match [batch, latent_length, 256]"
262 );
263 let min_samples = ((segment_duration_seconds * 12.5) as usize).max(1);
264 let mut hop_samples = (min_samples / HEARTCODEC_WINDOW_FRAMES) * 80;
265 let mut ovlp_samples = min_samples.saturating_sub(hop_samples);
266 if hop_samples == 0 {
267 hop_samples = 1;
268 ovlp_samples = min_samples;
269 }
270 let ovlp_frames = ovlp_samples * 2;
271 let target_len = (duration_seconds * HEARTMULA_SAMPLE_RATE as f32) as usize;
272 let audio_target_len = target_len;
273 let mut codes = codes;
274 if seq_len < min_samples {
275 while codes.dims()[2] < min_samples {
276 codes = Tensor::cat(vec![codes.clone(), codes], 2);
277 }
278 codes = codes.slice([0..batch, 0..num_quantizers, 0..min_samples]);
279 }
280
281 let mut codes_len = codes.dims()[2];
282 if !(codes_len.saturating_sub(ovlp_frames)).is_multiple_of(hop_samples) {
283 let len_codes = codes_len.saturating_sub(ovlp_samples).div_ceil(hop_samples)
284 * hop_samples
285 + ovlp_samples;
286 while codes.dims()[2] < len_codes {
287 codes = Tensor::cat(vec![codes.clone(), codes], 2);
288 }
289 codes = codes.slice([0..batch, 0..num_quantizers, 0..len_codes]);
290 codes_len = len_codes;
291 }
292
293 let mut windows = Vec::new();
294 let mut previous_latent: Option<Tensor<B, 3>> = None;
295 let mut sinx = 0usize;
296 while sinx + min_samples <= codes_len {
297 let window_end = (sinx + min_samples).min(codes_len);
298 let codes_input = codes
299 .clone()
300 .slice([0..batch, 0..num_quantizers, sinx..window_end]);
301 let window_latent_length = codes_input.dims()[2] * 2;
302
303 if sinx == 0 || ovlp_frames == 0 {
304 let initial_window_latent = first_latent.clone().slice([
305 0..batch,
306 0..window_latent_length.min(first_latent.dims()[1]),
307 0..first_latent.dims()[2],
308 ]);
309 let latent = Self::run_flow_matching_window(
310 flow_matching,
311 guidance_scale,
312 ode_steps,
313 codes_input,
314 initial_window_latent.clone(),
315 window_latent_length,
316 0,
317 Some(initial_window_latent),
318 );
319 windows.push(Self::latent_to_scalar_input(latent.clone()));
320 previous_latent = Some(latent);
321 } else {
322 let prev_latent = previous_latent
323 .as_ref()
324 .expect("previous latent window is required when overlap is enabled");
325 let true_latent = prev_latent.clone().slice([
326 0..batch,
327 prev_latent.dims()[1].saturating_sub(ovlp_frames)..prev_latent.dims()[1],
328 0..prev_latent.dims()[2],
329 ]);
330 let len_add_to_latent = window_latent_length.saturating_sub(true_latent.dims()[1]);
331 let true_latent = if len_add_to_latent == 0 {
332 true_latent
333 } else {
334 Tensor::cat(
335 vec![
336 true_latent.clone(),
337 Tensor::<B, 3>::random(
338 [batch, len_add_to_latent, true_latent.dims()[2]],
339 burn::tensor::Distribution::Normal(0.0, 1.0),
340 &device,
341 ),
342 ],
343 1,
344 )
345 };
346 let latent = Self::run_flow_matching_window(
347 flow_matching,
348 guidance_scale,
349 ode_steps,
350 codes_input,
351 true_latent,
352 window_latent_length,
353 ovlp_frames,
354 None,
355 );
356 windows.push(Self::latent_to_scalar_input(latent.clone()));
357 previous_latent = Some(latent);
358 }
359 sinx += hop_samples.max(1);
360 }
361
362 ScalarDecodePlan {
363 target_len,
364 audio_target_len,
365 windows,
366 }
367 }
368
369 pub(crate) fn decode_scalar_plan_impl(
370 scalar_model: &ScalarModel<B>,
371 plan: ScalarDecodePlan<B>,
372 ) -> Tensor<B, 3> {
373 let device = plan
374 .windows
375 .first()
376 .expect("expected at least one scalar decode window")
377 .device();
378 let min_samples = plan
379 .windows
380 .first()
381 .map(|window| window.dims()[2] * HEARTMULA_SAMPLE_RATE / 25)
382 .expect("expected at least one scalar decode window");
383 let hop_samples = ((min_samples / HEARTCODEC_WINDOW_FRAMES) * 80).max(1);
384 let ovlp_samples = min_samples.saturating_sub(hop_samples);
385 let mut output: Option<Tensor<B, 3>> = None;
386
387 for scalar_input in plan.windows {
388 let mut cur_output = scalar_model.decode_with_sync(scalar_input);
389 let cur_output_dims = cur_output.dims();
390 cur_output = cur_output.slice([
391 0..cur_output_dims[0],
392 0..1,
393 0..min_samples.min(cur_output_dims[2]),
394 ]);
395 if let Some(prev) = output {
396 if ovlp_samples == 0 {
397 output = Some(Tensor::cat(vec![prev, cur_output], 2));
398 } else {
399 let ov_win = {
400 let mut v = Vec::with_capacity(ovlp_samples);
401 for i in 0..ovlp_samples {
402 let denom = (ovlp_samples.saturating_sub(1)).max(1) as f32;
403 v.push(i as f32 / denom);
404 }
405 Tensor::<B, 3>::from_data(TensorData::new(v, [1, 1, ovlp_samples]), &device)
406 };
407 let prev_dims = prev.dims();
408 let prev_len = prev_dims[2];
409 let prev_head =
410 prev.clone()
411 .slice([0..prev_dims[0], 0..1, 0..prev_len - ovlp_samples]);
412 let prev_tail =
413 prev.slice([0..prev_dims[0], 0..1, prev_len - ovlp_samples..prev_len]);
414 let cur_dims = cur_output.dims();
415 let cur_head =
416 cur_output
417 .clone()
418 .slice([0..cur_dims[0], 0..1, 0..ovlp_samples]);
419 let prev_energy = prev_tail.clone().square();
420 let cur_energy = cur_head.clone().square();
421 let energy_sum = prev_energy.clone() + cur_energy.clone() + 1.0e-8;
422 let transient_cur = cur_energy / energy_sum;
423 let cur_weight = (ov_win.clone() + transient_cur) * 0.5;
424 let prev_weight =
425 Tensor::<B, 3>::ones([1, 1, ovlp_samples], &device) - cur_weight.clone();
426 let blended = prev_tail * prev_weight + cur_head * cur_weight;
427 output = Some(Tensor::cat(
428 vec![
429 prev_head,
430 blended,
431 cur_output.slice([0..cur_dims[0], 0..1, ovlp_samples..cur_dims[2]]),
432 ],
433 2,
434 ));
435 }
436 } else {
437 output = Some(cur_output);
438 }
439 }
440
441 output
442 .expect("expected at least one decoded window")
443 .slice([0..2, 0..1, 0..plan.target_len])
444 }
445
446 #[allow(clippy::too_many_arguments)]
447 fn run_flow_matching_window(
448 flow_matching: &FlowMatching<B>,
449 guidance_scale: f32,
450 ode_steps: usize,
451 codes: Tensor<B, 3, Int>,
452 true_latents: Tensor<B, 3>,
453 latent_length: usize,
454 incontext_length: usize,
455 initial_latent_override: Option<Tensor<B, 3>>,
456 ) -> Tensor<B, 3> {
457 flow_matching.inference_codes(
458 vec![codes],
459 true_latents,
460 latent_length,
461 incontext_length,
462 guidance_scale,
463 ode_steps,
464 false,
465 "other_seg",
466 initial_latent_override,
467 )
468 }
469
470 fn latent_to_scalar_input(latent: Tensor<B, 3>) -> Tensor<B, 3> {
471 let [batch, seq_len, channels] = latent.dims();
472 assert_eq!(channels, 256, "Expected 256 channels from flow matching");
473
474 let latent_reshaped = latent.reshape([batch, seq_len, 2, 128]);
475 let latent_permuted = latent_reshaped.swap_dims(1, 2);
476 let latent_split = latent_permuted.reshape([batch * 2, seq_len, 128]);
477
478 latent_split.swap_dims(1, 2)
479 }
480
481 pub fn build_scalar_decode_plan(
482 &self,
483 codes: Tensor<B, 3, Int>,
484 first_latent: Tensor<B, 3>,
485 ) -> ScalarDecodePlan<B> {
486 Self::build_scalar_decode_plan_impl(
487 &self.flow_matching,
488 self.guidance_scale,
489 self.ode_steps,
490 codes,
491 first_latent,
492 )
493 }
494
495 pub fn decode_scalar_plan(&self, plan: ScalarDecodePlan<B>) -> Tensor<B, 3> {
496 Self::decode_scalar_plan_impl(&self.scalar_model, plan)
497 }
498
499 pub fn decode(&self, codes: Tensor<B, 3, Int>) -> Tensor<B, 3> {
500 let [batch, _num_quantizers, _seq_len] = codes.dims();
501 let latent_length = (HEARTCODEC_SEGMENT_DURATION_SECONDS * 25.0) as usize;
502 let first_latent = Tensor::<B, 3>::random(
503 [batch, latent_length, 256],
504 burn::tensor::Distribution::Normal(0.0, 1.0),
505 &codes.device(),
506 );
507 self.decode_with_initial_latent(codes, first_latent)
508 }
509
510 pub fn decode_with_initial_latent(
511 &self,
512 codes: Tensor<B, 3, Int>,
513 first_latent: Tensor<B, 3>,
514 ) -> Tensor<B, 3> {
515 let plan = self.build_scalar_decode_plan(codes, first_latent);
516 self.decode_scalar_plan(plan)
517 }
518}
519
520#[derive(Module, Debug)]
521pub struct FlowMatching<B: Backend> {
522 pub cond_feature_emb: Linear<B>,
523 pub zero_cond_embedding1: Param<Tensor<B, 1>>,
524 pub estimator: LlamaTransformer<B>,
525 pub vq_embed: ResidualVQ<B>,
526 pub debug_latent_steps: bool,
527}
528
529impl<B: Backend> FlowMatching<B> {
530 pub fn new(device: &B::Device, config: &HeartCodecConfig) -> Self {
531 Self {
532 cond_feature_emb: LinearConfig::new(config.dim, config.dim)
533 .with_bias(true)
534 .with_layout(LinearLayout::Col)
535 .init(device),
536 zero_cond_embedding1: Param::from_tensor(Tensor::zeros([config.dim], device)),
537 estimator: LlamaTransformer::new(device, config),
538 vq_embed: ResidualVQ::new(device, config),
539 debug_latent_steps: false,
540 }
541 }
542
543 pub fn load_from_burnpack(path: &std::path::Path, device: &B::Device) -> Result<Self> {
544 let mut model = Self::new(device, &HeartCodecConfig::default());
545 let mut store = BurnpackStore::from_file(path).zero_copy(true);
546
547 if let Err(load_err) = model.load_from(&mut store) {
548 let get_tensor = HeartCodecModel::<B>::snapshots_to_f32_lookup(path)?;
549 return HeartCodecModel::<B>::load_flow_matching_manually(
550 &mut model, path, device, &get_tensor,
551 )
552 .map_err(|fallback_err| {
553 anyhow::anyhow!(
554 "Failed to load FlowMatching: {load_err}; fallback mapping also failed: {fallback_err}"
555 )
556 });
557 }
558
559 Ok(model)
560 }
561
562 fn interpolate_1d(x: &Tensor<B, 3>, scale_factor: usize) -> Tensor<B, 3> {
563 let [batch, seq_len, channels] = x.dims();
564 if scale_factor <= 1 {
565 return x.clone();
566 }
567
568 let mut repeated_steps = Vec::with_capacity(seq_len * scale_factor);
569 for step in 0..seq_len {
570 let slice = x.clone().slice([0..batch, step..step + 1, 0..channels]);
571 for _ in 0..scale_factor {
572 repeated_steps.push(slice.clone());
573 }
574 }
575 Tensor::cat(repeated_steps, 1)
576 }
577
578 #[allow(clippy::too_many_arguments)]
579 pub fn solve_ode(
580 &self,
581 conditioning: Tensor<B, 3>,
582 true_latents: Tensor<B, 3>,
583 _latent_length: usize,
584 incontext_length: usize,
585 num_steps: usize,
586 guidance_scale: f32,
587 initial_latent_override: Option<Tensor<B, 3>>,
588 ) -> Tensor<B, 3> {
589 let device = conditioning.device();
590 let [batch, _seq_len, _cond_dim] = conditioning.dims();
591
592 let latent_dim = 256;
593 let num_steps = num_steps.clamp(1, 50);
594
595 let cond_interp = Self::interpolate_1d(&conditioning, 2);
596 let [_batch, seq_len_interp, _] = cond_interp.dims();
597 let latent_masks =
598 Self::build_latent_masks(seq_len_interp, _latent_length, incontext_length);
599 let masked_incontext_length = latent_masks.iter().filter(|&&mask| mask == 1).count();
600 let zero_cond = self
601 .zero_cond_embedding1
602 .val()
603 .clone()
604 .reshape([1, 1, 512])
605 .repeat(&[batch, seq_len_interp, 1]);
606 let active_mask = Tensor::<B, 3>::from_data(
607 TensorData::new(
608 latent_masks
609 .iter()
610 .map(|&mask| if mask > 0 { 1.0 } else { 0.0 })
611 .collect::<Vec<_>>(),
612 [1, seq_len_interp, 1],
613 ),
614 &device,
615 )
616 .repeat(&[batch, 1, 512]);
617 let inactive_mask =
618 Tensor::<B, 3>::ones([batch, seq_len_interp, 512], &device) - active_mask.clone();
619 let cond_with_mask = cond_interp.clone() * active_mask + zero_cond.clone() * inactive_mask;
620 let uncond_mask = Tensor::<B, 3>::zeros([batch, seq_len_interp, 512], &device);
621
622 let mut latent = if let Some(initial_latent) = initial_latent_override {
623 assert_eq!(
624 initial_latent.dims(),
625 [batch, seq_len_interp, latent_dim],
626 "initial latent override shape must match [batch, seq_len_interp, latent_dim]"
627 );
628 initial_latent
629 } else {
630 Tensor::<B, 3>::random(
631 [batch, seq_len_interp, latent_dim],
632 burn::tensor::Distribution::Normal(0.0, 1.0),
633 &device,
634 )
635 };
636 let incontext_mask = latent_masks
637 .iter()
638 .map(|&m| if m == 1 { 1_i64 } else { 0_i64 })
639 .collect::<Vec<_>>();
640 let incontext_mask = Tensor::<B, 3>::from_data(
641 TensorData::new(
642 incontext_mask
643 .into_iter()
644 .map(|v| v as f32)
645 .collect::<Vec<_>>(),
646 [1, seq_len_interp, 1],
647 ),
648 &device,
649 );
650 let incontext_x = true_latents * incontext_mask;
651
652 let dt = 1.0 / num_steps as f32;
653
654 let sync_interval = if num_steps <= 5 { 5 } else { 3 };
655 for step in 0..num_steps {
656 let t = step as f32 * dt;
657
658 if masked_incontext_length > 0 {
659 let noise = latent.clone();
660 let prefix =
661 noise
662 .clone()
663 .slice([0..batch, 0..masked_incontext_length, 0..latent_dim]);
664 let incontext_prefix = incontext_x.clone().slice([
665 0..batch,
666 0..masked_incontext_length,
667 0..latent_dim,
668 ]);
669 let anchored = prefix * (1.0 - (1.0 - 1e-6) * t) + incontext_prefix * t;
670 let suffix = noise.slice([
671 0..batch,
672 masked_incontext_length..seq_len_interp,
673 0..latent_dim,
674 ]);
675 latent = Tensor::cat(vec![anchored, suffix], 1);
676 }
677
678 let velocity = if guidance_scale > 1.0 {
679 let uncond_input = Tensor::cat(
680 vec![latent.clone(), incontext_x.clone(), uncond_mask.clone()],
681 2,
682 );
683 let uncond_vel = self.estimator.forward(&uncond_input, t, step);
684
685 let cond_input = Tensor::cat(
686 vec![latent.clone(), incontext_x.clone(), cond_with_mask.clone()],
687 2,
688 );
689 let cond_vel = self.estimator.forward(&cond_input, t, step);
690
691 uncond_vel.clone() + (cond_vel - uncond_vel) * guidance_scale
692 } else {
693 let estimator_input = Tensor::cat(
694 vec![latent.clone(), incontext_x.clone(), cond_with_mask.clone()],
695 2,
696 );
697 self.estimator.forward(&estimator_input, t, step)
698 };
699
700 latent = latent + velocity * dt;
701
702 if step > 0 && step % sync_interval == 0 {
703 let _ = latent.to_data();
704 }
705 }
706
707 let _ = latent.to_data();
708
709 if masked_incontext_length > 0 {
710 let prefix =
711 incontext_x
712 .clone()
713 .slice([0..batch, 0..masked_incontext_length, 0..latent_dim]);
714 let suffix = latent.slice([
715 0..batch,
716 masked_incontext_length..seq_len_interp,
717 0..latent_dim,
718 ]);
719 latent = Tensor::cat(vec![prefix, suffix], 1);
720 }
721
722 latent
723 }
724
725 #[allow(clippy::too_many_arguments)]
726 pub fn inference_codes(
727 &self,
728 codes: Vec<Tensor<B, 3, Int>>,
729 true_latents: Tensor<B, 3>,
730 latent_length: usize,
731 incontext_length: usize,
732 guidance_scale: f32,
733 num_steps: usize,
734 disable_progress: bool,
735 scenario: &str,
736 initial_latent_override: Option<Tensor<B, 3>>,
737 ) -> Tensor<B, 3> {
738 let _ = disable_progress;
739 let codes_bestrq_emb = codes
740 .into_iter()
741 .next()
742 .expect("inference_codes expects at least one codes tensor");
743 let conditioning = self.get_output_from_indices(codes_bestrq_emb);
744 let conditioning = self.cond_feature_emb.forward(conditioning);
745 let _ = scenario;
746 self.solve_ode(
747 conditioning,
748 true_latents,
749 latent_length,
750 incontext_length,
751 num_steps,
752 guidance_scale,
753 initial_latent_override,
754 )
755 }
756
757 fn gather_codebook(
758 &self,
759 embed: Tensor<B, 3>,
760 indices: Tensor<B, 2, Int>,
761 codebook_size: usize,
762 dim: usize,
763 ) -> Tensor<B, 3> {
764 let [batch, seq_len] = indices.dims();
765 let embed_2d: Tensor<B, 2> = embed.squeeze_dim(0);
766 let indices_flat: Tensor<B, 1, Int> = indices.reshape([batch * seq_len]);
767 let max_idx = (codebook_size as i64) - 1;
768 let indices_clamped: Tensor<B, 1, Int> = indices_flat.clamp(0, max_idx);
769 let gathered = embed_2d.select(0, indices_clamped);
770 gathered.reshape([batch, seq_len, dim])
771 }
772
773 fn get_output_from_indices(&self, codes: Tensor<B, 3, Int>) -> Tensor<B, 3> {
774 let [batch, num_quantizers, seq_len] = codes.dims();
775
776 let mut quantized_sum = Tensor::<B, 3>::zeros([batch, seq_len, 32], &codes.device());
777 for q in 0..num_quantizers.min(self.vq_embed.layers.len()) {
778 let q_codes_3d = codes.clone().slice([0..batch, q..q + 1, 0..seq_len]);
779 let q_codes = q_codes_3d.reshape([batch, seq_len]);
780 let embed = self.vq_embed.layers[q]._codebook.embed.val();
781 let embed_dim = embed.dims()[2];
782 let codebook_size = embed.dims()[1];
783 let q_emb = self.gather_codebook(embed, q_codes, codebook_size, embed_dim);
784 quantized_sum = quantized_sum + q_emb;
785 }
786 self.vq_embed.project_out.forward(quantized_sum)
787 }
788
789 fn build_latent_masks(
790 seq_len: usize,
791 latent_length: usize,
792 incontext_length: usize,
793 ) -> Vec<i64> {
794 let mut masks = vec![0_i64; seq_len];
795 for mask in masks.iter_mut().take(seq_len.min(latent_length)) {
796 *mask = 2;
797 }
798 for mask in masks.iter_mut().take(seq_len.min(incontext_length)) {
799 *mask = 1;
800 }
801 masks
802 }
803}
804
805#[derive(Module, Debug)]
806pub struct ResidualVQ<B: Backend> {
807 pub layers: Vec<VQCodebook<B>>,
808 pub project_in: Linear<B>,
809 pub project_out: Linear<B>,
810}
811
812impl<B: Backend> ResidualVQ<B> {
813 pub fn new(device: &B::Device, config: &HeartCodecConfig) -> Self {
814 let layers: Vec<_> = (0..config.num_quantizers)
815 .map(|_| VQCodebook::new(device, config.codebook_size, config.codebook_dim))
816 .collect();
817
818 Self {
819 layers,
820 project_in: LinearConfig::new(512, 32)
821 .with_bias(true)
822 .with_layout(LinearLayout::Col)
823 .init(device),
824 project_out: LinearConfig::new(32, 512)
825 .with_bias(true)
826 .with_layout(LinearLayout::Col)
827 .init(device),
828 }
829 }
830
831 pub fn load_from_dot_notation<F>(device: &B::Device, get_tensor: &F) -> Result<Self>
832 where
833 F: Fn(&str) -> Option<(Vec<f32>, Vec<usize>)>,
834 {
835 let mut layers = Vec::new();
836 for i in 0..8 {
837 let prefix = format!("flow_matching.vq_embed.layers.{}._codebook", i);
838
839 let embed_name = format!("{}.embed", prefix);
840 let embed = if let Some((data, shape)) = get_tensor(&embed_name) {
841 Tensor::<B, 3>::from_data(
842 TensorData::new(data, [shape[0], shape[1], shape[2]]),
843 device,
844 )
845 } else {
846 Tensor::zeros([1, 8192, 32], device)
847 };
848
849 let cluster_size_name = format!("{}.cluster_size", prefix);
850 let cluster_size = if let Some((data, shape)) = get_tensor(&cluster_size_name) {
851 Tensor::<B, 2>::from_data(TensorData::new(data, [shape[0], shape[1]]), device)
852 } else {
853 Tensor::zeros([1, 8192], device)
854 };
855
856 let embed_avg_name = format!("{}.embed_avg", prefix);
857 let embed_avg = if let Some((data, shape)) = get_tensor(&embed_avg_name) {
858 Tensor::<B, 3>::from_data(
859 TensorData::new(data, [shape[0], shape[1], shape[2]]),
860 device,
861 )
862 } else {
863 Tensor::zeros([1, 8192, 32], device)
864 };
865
866 layers.push(VQCodebook {
867 _codebook: VQCodebookInner {
868 embed: Param::from_tensor(embed),
869 cluster_size: Param::from_tensor(cluster_size),
870 embed_avg: Param::from_tensor(embed_avg),
871 },
872 });
873 }
874
875 let project_in = load_linear_from_tensors(
876 device,
877 get_tensor,
878 "flow_matching.vq_embed.project_in",
879 512,
880 32,
881 )?;
882
883 let project_out = load_linear_from_tensors(
884 device,
885 get_tensor,
886 "flow_matching.vq_embed.project_out",
887 32,
888 512,
889 )?;
890
891 Ok(Self {
892 layers,
893 project_in,
894 project_out,
895 })
896 }
897}
898
899#[derive(Module, Debug)]
900pub struct VQCodebook<B: Backend> {
901 pub _codebook: VQCodebookInner<B>,
902}
903
904impl<B: Backend> VQCodebook<B> {
905 pub fn new(device: &B::Device, codebook_size: usize, codebook_dim: usize) -> Self {
906 Self {
907 _codebook: VQCodebookInner::new(device, codebook_size, codebook_dim),
908 }
909 }
910}
911
912#[derive(Module, Debug)]
913pub struct VQCodebookInner<B: Backend> {
914 pub cluster_size: Param<Tensor<B, 2>>,
915 pub embed: Param<Tensor<B, 3>>,
916 pub embed_avg: Param<Tensor<B, 3>>,
917}
918
919impl<B: Backend> VQCodebookInner<B> {
920 pub fn new(device: &B::Device, codebook_size: usize, codebook_dim: usize) -> Self {
921 Self {
922 cluster_size: Param::from_tensor(Tensor::zeros([1, codebook_size], device)),
923 embed: Param::from_tensor(Tensor::zeros([1, codebook_size, codebook_dim], device)),
924 embed_avg: Param::from_tensor(Tensor::zeros([1, codebook_size, codebook_dim], device)),
925 }
926 }
927}
928
929#[derive(Module, Debug)]
930pub struct LlamaTransformer<B: Backend> {
931 pub proj_in: ProjectLayer<B>,
932 pub proj_out: ProjectLayer<B>,
933 pub connection_proj: ProjectLayer<B>,
934 pub transformer_blocks: Vec<TransformerBlock<B>>,
935 pub transformer_blocks_2: Vec<TransformerBlock<B>>,
936 pub norm_out: LayerNorm<B>,
937 pub norm_out_2: LayerNorm<B>,
938 pub adaln_single: AdaLayerNormSingle<B>,
939 pub adaln_single_2: AdaLayerNormSingle<B>,
940 pub scale_shift_table: Param<Tensor<B, 2>>,
941 pub scale_shift_table_2: Param<Tensor<B, 2>>,
942}
943
944impl<B: Backend> LlamaTransformer<B> {
945 pub fn new(device: &B::Device, config: &HeartCodecConfig) -> Self {
946 let inner_dim = config.num_attention_heads * config.attention_head_dim;
947 let inner_dim_2 = inner_dim * 2;
948 let _latent_dim = config.latent_hidden_dim;
949
950 let transformer_blocks: Vec<_> = (0..config.num_layers)
951 .map(|_| {
952 TransformerBlock::new(
953 device,
954 inner_dim,
955 config.num_attention_heads,
956 config.attention_head_dim,
957 )
958 })
959 .collect();
960
961 let transformer_blocks_2: Vec<_> = (0..config.num_layers_2)
962 .map(|_| {
963 TransformerBlock::new(
964 device,
965 inner_dim_2,
966 config.num_attention_heads,
967 config.attention_head_dim * 2,
968 )
969 })
970 .collect();
971
972 let in_channels = 1024;
973 let connection_in = 2560;
974
975 Self {
976 proj_in: ProjectLayer::new(device, in_channels, inner_dim, 3),
977 proj_out: ProjectLayer::new(device, inner_dim_2, config.out_channels, 3),
978 connection_proj: ProjectLayer::new(device, connection_in, inner_dim_2, 3),
979 transformer_blocks,
980 transformer_blocks_2,
981 norm_out: LayerNormConfig::new(inner_dim)
982 .with_epsilon(1e-6)
983 .with_bias(false)
984 .init(device),
985 norm_out_2: LayerNormConfig::new(inner_dim_2)
986 .with_epsilon(1e-6)
987 .with_bias(false)
988 .init(device),
989 adaln_single: AdaLayerNormSingle::new(device, inner_dim),
990 adaln_single_2: AdaLayerNormSingle::new(device, inner_dim_2),
991 scale_shift_table: Param::from_tensor(Tensor::zeros([2, inner_dim], device)),
992 scale_shift_table_2: Param::from_tensor(Tensor::zeros([2, inner_dim_2], device)),
993 }
994 }
995
996 pub fn load_from_burnpack<F>(device: &B::Device, get_tensor: &F) -> Result<Self>
997 where
998 F: Fn(&str) -> Option<(Vec<f32>, Vec<usize>)>,
999 {
1000 let config = HeartCodecConfig::default();
1001 let inner_dim = config.num_attention_heads * config.attention_head_dim;
1002 let inner_dim_2 = inner_dim * 2;
1003 let in_channels = 1024;
1004 let connection_in = 2560;
1005
1006 let proj_in = ProjectLayer::load_from_tensors(
1007 device,
1008 get_tensor,
1009 "flow_matching.estimator.proj_in",
1010 in_channels,
1011 inner_dim,
1012 )?;
1013
1014 let proj_out = ProjectLayer::load_from_tensors(
1015 device,
1016 get_tensor,
1017 "flow_matching.estimator.proj_out",
1018 inner_dim_2,
1019 config.out_channels,
1020 )?;
1021
1022 let connection_proj = ProjectLayer::load_from_tensors(
1023 device,
1024 get_tensor,
1025 "flow_matching.estimator.connection_proj",
1026 connection_in,
1027 inner_dim_2,
1028 )?;
1029
1030 let mut transformer_blocks = Vec::new();
1031 for i in 0..config.num_layers {
1032 let block = TransformerBlock::load_from_tensors(
1033 device,
1034 get_tensor,
1035 &format!("flow_matching.estimator.transformer_blocks.{}", i),
1036 inner_dim,
1037 config.num_attention_heads,
1038 config.attention_head_dim,
1039 )?;
1040 transformer_blocks.push(block);
1041 }
1042
1043 let mut transformer_blocks_2 = Vec::new();
1044 for i in 0..config.num_layers_2 {
1045 let block = TransformerBlock::load_from_tensors(
1046 device,
1047 get_tensor,
1048 &format!("flow_matching.estimator.transformer_blocks_2.{}", i),
1049 inner_dim_2,
1050 config.num_attention_heads,
1051 config.attention_head_dim * 2,
1052 )?;
1053 transformer_blocks_2.push(block);
1054 }
1055
1056 let adaln_single = AdaLayerNormSingle::load_from_tensors(
1057 device,
1058 get_tensor,
1059 "flow_matching.estimator.adaln_single",
1060 inner_dim,
1061 )?;
1062
1063 let adaln_single_2 = AdaLayerNormSingle::load_from_tensors(
1064 device,
1065 get_tensor,
1066 "flow_matching.estimator.adaln_single_2",
1067 inner_dim_2,
1068 )?;
1069
1070 let scale_shift_table = load_param_tensor(
1071 device,
1072 get_tensor,
1073 "flow_matching.estimator.scale_shift_table",
1074 [2, inner_dim],
1075 )?;
1076
1077 let scale_shift_table_2 = load_param_tensor(
1078 device,
1079 get_tensor,
1080 "flow_matching.estimator.scale_shift_table_2",
1081 [2, inner_dim_2],
1082 )?;
1083
1084 Ok(Self {
1085 proj_in,
1086 proj_out,
1087 connection_proj,
1088 transformer_blocks,
1089 transformer_blocks_2,
1090 norm_out: LayerNormConfig::new(inner_dim)
1091 .with_epsilon(1e-6)
1092 .with_bias(false)
1093 .init(device),
1094 norm_out_2: LayerNormConfig::new(inner_dim_2)
1095 .with_epsilon(1e-6)
1096 .with_bias(false)
1097 .init(device),
1098 adaln_single,
1099 adaln_single_2,
1100 scale_shift_table,
1101 scale_shift_table_2,
1102 })
1103 }
1104
1105 pub fn forward(&self, hidden_states: &Tensor<B, 3>, t: f32, step: usize) -> Tensor<B, 3> {
1106 let mut s = self.proj_in.forward(hidden_states.clone(), step);
1107 let (timestep_mod, embedded_timestep) = self.adaln_single.forward(t, s.dtype());
1108
1109 for block in &self.transformer_blocks {
1110 s = block.forward(s, Some(timestep_mod.clone()), false, step);
1111 }
1112
1113 let shift_scale_1 = self.scale_shift_table.val().clone().unsqueeze_dim(0)
1114 + embedded_timestep.unsqueeze_dim(1);
1115 let shift_1 = shift_scale_1.clone().slice([0..1, 0..1, 0..s.dims()[2]]);
1116 let scale_1 = shift_scale_1.slice([0..1, 1..2, 0..s.dims()[2]]);
1117 let s_norm = self.norm_out.forward(s);
1118 let s = s_norm * (scale_1 + 1.0) + shift_1;
1119
1120 let x = Tensor::cat(vec![hidden_states.clone(), s.clone()], 2);
1121
1122 let x = self.connection_proj.forward(x, step);
1123
1124 let mut x = x;
1125 let (timestep_mod_2, embedded_timestep_2) = self.adaln_single_2.forward(t, x.dtype());
1126 for block in &self.transformer_blocks_2 {
1127 x = block.forward(x, Some(timestep_mod_2.clone()), false, step);
1128 }
1129
1130 let shift_scale_2 = self.scale_shift_table_2.val().clone().unsqueeze_dim(0)
1131 + embedded_timestep_2.unsqueeze_dim(1);
1132 let shift_2 = shift_scale_2.clone().slice([0..1, 0..1, 0..x.dims()[2]]);
1133 let scale_2 = shift_scale_2.slice([0..1, 1..2, 0..x.dims()[2]]);
1134 let x_norm = self.norm_out_2.forward(x);
1135 let x = x_norm * (scale_2 + 1.0) + shift_2;
1136
1137 self.proj_out.forward(x, step)
1138 }
1139}
1140
1141#[derive(Module, Debug)]
1142pub struct TransformerBlock<B: Backend> {
1143 pub attn: Attention<B>,
1144 pub attn_norm: RmsNorm<B>,
1145 pub mlp: Mlp<B>,
1146 pub mlp_norm: RmsNorm<B>,
1147 pub scale_shift_table: Param<Tensor<B, 2>>,
1148}
1149
1150impl<B: Backend> TransformerBlock<B> {
1151 pub fn new(device: &B::Device, dim: usize, num_heads: usize, head_dim: usize) -> Self {
1152 Self {
1153 attn: Attention::new(device, dim, num_heads, head_dim),
1154 attn_norm: RmsNorm::new(device, dim, 1e-6),
1155 mlp: Mlp::new(device, dim),
1156 mlp_norm: RmsNorm::new(device, dim, 1e-6),
1157 scale_shift_table: Param::from_tensor(Tensor::zeros([6, dim], device)),
1158 }
1159 }
1160
1161 pub fn load_from_tensors<F>(
1162 device: &B::Device,
1163 get_tensor: &F,
1164 prefix: &str,
1165 dim: usize,
1166 num_heads: usize,
1167 head_dim: usize,
1168 ) -> Result<Self>
1169 where
1170 F: Fn(&str) -> Option<(Vec<f32>, Vec<usize>)>,
1171 {
1172 let inner_dim = num_heads * head_dim;
1173 let hidden_dim = Mlp::<B>::compute_hidden_dim(dim);
1174
1175 let attn = Attention {
1176 q_proj: load_linear_from_tensors(
1177 device,
1178 get_tensor,
1179 &format!("{}.attn.q_proj", prefix),
1180 dim,
1181 inner_dim,
1182 )?,
1183 k_proj: load_linear_from_tensors(
1184 device,
1185 get_tensor,
1186 &format!("{}.attn.k_proj", prefix),
1187 dim,
1188 inner_dim,
1189 )?,
1190 v_proj: load_linear_from_tensors(
1191 device,
1192 get_tensor,
1193 &format!("{}.attn.v_proj", prefix),
1194 dim,
1195 inner_dim,
1196 )?,
1197 o_proj: load_linear_from_tensors(
1198 device,
1199 get_tensor,
1200 &format!("{}.attn.o_proj", prefix),
1201 inner_dim,
1202 dim,
1203 )?,
1204 num_heads,
1205 head_dim,
1206 rope_dim: head_dim,
1207 };
1208
1209 let mlp = Mlp {
1210 gate: load_linear_from_tensors(
1211 device,
1212 get_tensor,
1213 &format!("{}.mlp.gate", prefix),
1214 dim,
1215 hidden_dim,
1216 )?,
1217 up: load_linear_from_tensors(
1218 device,
1219 get_tensor,
1220 &format!("{}.mlp.up", prefix),
1221 dim,
1222 hidden_dim,
1223 )?,
1224 down: load_linear_from_tensors(
1225 device,
1226 get_tensor,
1227 &format!("{}.mlp.down", prefix),
1228 hidden_dim,
1229 dim,
1230 )?,
1231 };
1232
1233 let attn_norm =
1234 load_rmsnorm_from_tensors(device, get_tensor, &format!("{}.attn_norm", prefix), dim)?;
1235
1236 let mlp_norm =
1237 load_rmsnorm_from_tensors(device, get_tensor, &format!("{}.mlp_norm", prefix), dim)?;
1238
1239 let scale_shift_table = load_param_tensor(
1240 device,
1241 get_tensor,
1242 &format!("{}.scale_shift_table", prefix),
1243 [6, dim],
1244 )?;
1245
1246 Ok(Self {
1247 attn,
1248 attn_norm,
1249 mlp,
1250 mlp_norm,
1251 scale_shift_table,
1252 })
1253 }
1254
1255 pub fn forward(
1256 &self,
1257 x: Tensor<B, 3>,
1258 timestep: Option<Tensor<B, 3>>,
1259 _dump_attention: bool,
1260 _step: usize,
1261 ) -> Tensor<B, 3> {
1262 if let Some(timestep) = timestep {
1263 let [batch, _seq, dim] = x.dims();
1264 let shift_scale = self.scale_shift_table.val().clone().unsqueeze_dim(0)
1265 + timestep.reshape([batch, 6, dim]);
1266 let mut parts = shift_scale.chunk(6, 1);
1267 let shift_msa = parts.remove(0);
1268 let scale_msa = parts.remove(0);
1269 let gate_msa = parts.remove(0);
1270 let shift_mlp = parts.remove(0);
1271 let scale_mlp = parts.remove(0);
1272 let gate_mlp = parts.remove(0);
1273
1274 let normed = self.attn_norm.forward(x.clone());
1275 let normed = normed * (scale_msa + 1.0) + shift_msa;
1276 let attn_out = self.attn.forward(normed, false, 0);
1277 let x = x + gate_msa * attn_out;
1278
1279 let normed = self.mlp_norm.forward(x.clone());
1280 let normed = normed * (scale_mlp + 1.0) + shift_mlp;
1281 let mlp_out = self.mlp.forward(normed);
1282 x + gate_mlp * mlp_out
1283 } else {
1284 let normed = self.attn_norm.forward(x.clone());
1285 let attn_out = self.attn.forward(normed, false, 0);
1286 let x = x + attn_out;
1287 let normed = self.mlp_norm.forward(x.clone());
1288 let mlp_out = self.mlp.forward(normed);
1289 x + mlp_out
1290 }
1291 }
1292}
1293
1294#[derive(Module, Debug)]
1295pub struct Attention<B: Backend> {
1296 pub q_proj: Linear<B>,
1297 pub k_proj: Linear<B>,
1298 pub v_proj: Linear<B>,
1299 pub o_proj: Linear<B>,
1300 pub num_heads: usize,
1301 pub head_dim: usize,
1302 pub rope_dim: usize,
1303}
1304
1305impl<B: Backend> Attention<B> {
1306 pub fn new(device: &B::Device, dim: usize, num_heads: usize, head_dim: usize) -> Self {
1307 let inner_dim = num_heads * head_dim;
1308 Self {
1309 q_proj: LinearConfig::new(dim, inner_dim)
1310 .with_bias(false)
1311 .with_layout(LinearLayout::Col)
1312 .init(device),
1313 k_proj: LinearConfig::new(dim, inner_dim)
1314 .with_bias(false)
1315 .with_layout(LinearLayout::Col)
1316 .init(device),
1317 v_proj: LinearConfig::new(dim, inner_dim)
1318 .with_bias(false)
1319 .with_layout(LinearLayout::Col)
1320 .init(device),
1321 o_proj: LinearConfig::new(inner_dim, dim)
1322 .with_bias(false)
1323 .with_layout(LinearLayout::Col)
1324 .init(device),
1325 num_heads,
1326 head_dim,
1327 rope_dim: head_dim,
1328 }
1329 }
1330
1331 pub fn forward(&self, x: Tensor<B, 3>, _dump_attention: bool, _step: usize) -> Tensor<B, 3> {
1332 let [batch, seq_len, dim] = x.dims();
1333 let num_heads = self.num_heads;
1334 let head_dim = self.head_dim;
1335
1336 let q = self.q_proj.forward(x.clone());
1337 let k = self.k_proj.forward(x.clone());
1338 let v = self.v_proj.forward(x);
1339
1340 let q = q
1341 .reshape([batch, seq_len, num_heads, head_dim])
1342 .swap_dims(1, 2);
1343 let k = k
1344 .reshape([batch, seq_len, num_heads, head_dim])
1345 .swap_dims(1, 2);
1346 let v = v
1347 .reshape([batch, seq_len, num_heads, head_dim])
1348 .swap_dims(1, 2);
1349
1350 let (q, k) = Self::apply_rope(q, k, self.rope_dim.min(head_dim));
1351
1352 let scores = q.matmul(k.swap_dims(2, 3)) / (head_dim as f32).sqrt();
1353
1354 use burn::tensor::activation::softmax;
1355 let attn_weights = softmax(scores, 3);
1356
1357 let out = attn_weights.matmul(v);
1358
1359 let out = out.swap_dims(1, 2).reshape([batch, seq_len, dim]);
1360
1361 self.o_proj.forward(out)
1362 }
1363
1364 fn apply_rope(
1365 q: Tensor<B, 4>,
1366 k: Tensor<B, 4>,
1367 rope_dim: usize,
1368 ) -> (Tensor<B, 4>, Tensor<B, 4>) {
1369 if rope_dim == 0 {
1370 return (q, k);
1371 }
1372
1373 let [batch, heads, seq_len, head_dim] = q.dims();
1374 let rope_pairs = rope_dim / 2;
1375 let device = q.device();
1376 let dtype = q.dtype();
1377 let base = 10_000.0_f32;
1378
1379 let mut inv_freq = Vec::with_capacity(rope_pairs);
1380 for i in 0..rope_pairs {
1381 let exponent = (2 * i) as f32 / rope_dim as f32;
1382 inv_freq.push(1.0 / base.powf(exponent));
1383 }
1384
1385 let inv_freq = Tensor::<B, 1>::from_data(TensorData::new(inv_freq, [rope_pairs]), &device);
1386 let positions = {
1387 let mut data = Vec::with_capacity(seq_len);
1388 for i in 0..seq_len {
1389 data.push(i as f32);
1390 }
1391 Tensor::<B, 2>::from_data(TensorData::new(data, [seq_len, 1]), &device).cast(dtype)
1392 };
1393 let freqs = positions.matmul(inv_freq.reshape([1, rope_pairs]));
1394 let sin = freqs.clone().sin().reshape([1, 1, seq_len, rope_pairs]);
1395 let cos = freqs.cos().reshape([1, 1, seq_len, rope_pairs]);
1396
1397 let rotate = |x: Tensor<B, 4>| {
1398 let head = x
1399 .clone()
1400 .slice([0..batch, 0..heads, 0..seq_len, 0..rope_dim]);
1401 let tail = x.slice([0..batch, 0..heads, 0..seq_len, rope_dim..head_dim]);
1402 let head = head.reshape([batch, heads, seq_len, rope_pairs, 2]);
1403 let x1 = head
1404 .clone()
1405 .slice([0..batch, 0..heads, 0..seq_len, 0..rope_pairs, 0..1])
1406 .reshape([batch, heads, seq_len, rope_pairs]);
1407 let x2 = head
1408 .slice([0..batch, 0..heads, 0..seq_len, 0..rope_pairs, 1..2])
1409 .reshape([batch, heads, seq_len, rope_pairs]);
1410 let rot_a = x1.clone() * cos.clone() - x2.clone() * sin.clone();
1411 let rot_b = x1 * sin.clone() + x2 * cos.clone();
1412 let rot = Tensor::cat(vec![rot_a, rot_b], 3);
1413 Tensor::cat(vec![rot, tail], 3)
1414 };
1415
1416 (rotate(q), rotate(k))
1417 }
1418}
1419
1420#[derive(Module, Debug)]
1421pub struct Mlp<B: Backend> {
1422 pub gate: Linear<B>,
1423 pub up: Linear<B>,
1424 pub down: Linear<B>,
1425}
1426
1427impl<B: Backend> Mlp<B> {
1428 pub fn new(device: &B::Device, dim: usize) -> Self {
1429 let hidden_dim = Self::compute_hidden_dim(dim);
1430 Self {
1431 gate: LinearConfig::new(dim, hidden_dim)
1432 .with_bias(false)
1433 .with_layout(LinearLayout::Col)
1434 .init(device),
1435 up: LinearConfig::new(dim, hidden_dim)
1436 .with_bias(false)
1437 .with_layout(LinearLayout::Col)
1438 .init(device),
1439 down: LinearConfig::new(hidden_dim, dim)
1440 .with_bias(false)
1441 .with_layout(LinearLayout::Col)
1442 .init(device),
1443 }
1444 }
1445
1446 fn compute_hidden_dim(dim: usize) -> usize {
1447 let multiple_of = 256;
1448 let hidden_dim = (4 * dim * 2) / 3;
1449 multiple_of * hidden_dim.div_ceil(multiple_of)
1450 }
1451
1452 pub fn forward(&self, x: Tensor<B, 3>) -> Tensor<B, 3> {
1453 use burn::tensor::activation::silu;
1454
1455 let gate = self.gate.forward(x.clone());
1456 let up = self.up.forward(x.clone());
1457
1458 let gate_activated = silu(gate);
1459
1460 let hidden = gate_activated * up;
1461
1462 self.down.forward(hidden)
1463 }
1464}
1465
1466#[derive(Module, Debug)]
1467pub struct RmsNorm<B: Backend> {
1468 pub weight: Param<Tensor<B, 1>>,
1469}
1470
1471impl<B: Backend> RmsNorm<B> {
1472 pub fn new(device: &B::Device, dim: usize, _eps: f64) -> Self {
1473 Self {
1474 weight: Param::from_tensor(Tensor::ones([dim], device)),
1475 }
1476 }
1477
1478 pub fn forward(&self, x: Tensor<B, 3>) -> Tensor<B, 3> {
1479 let eps = 1e-6;
1480 let weight = self.weight.val();
1481
1482 let x_sq = x.clone().powf_scalar(2.0);
1483
1484 let mean_sq = x_sq.mean_dim(2);
1485
1486 let rms = (mean_sq + eps).sqrt();
1487
1488 let [batch, seq_len, dim] = x.dims();
1489 let rms_expanded = rms.expand([batch, seq_len, dim]);
1490 let normalized = x / rms_expanded;
1491
1492 let weight_expanded = weight.reshape([1, 1, dim]).expand([batch, seq_len, dim]);
1493 normalized * weight_expanded
1494 }
1495}
1496
1497#[derive(Module, Debug)]
1498pub struct AdaLayerNormSingle<B: Backend> {
1499 pub emb: PixArtAlphaCombinedFlowEmbeddings<B>,
1500 pub linear: Linear<B>,
1501}
1502
1503impl<B: Backend> AdaLayerNormSingle<B> {
1504 pub fn new(device: &B::Device, embedding_dim: usize) -> Self {
1505 Self {
1506 emb: PixArtAlphaCombinedFlowEmbeddings::new(device, embedding_dim),
1507 linear: LinearConfig::new(embedding_dim, 6 * embedding_dim)
1508 .with_bias(true)
1509 .with_layout(LinearLayout::Col)
1510 .init(device),
1511 }
1512 }
1513
1514 pub fn load_from_tensors<F>(
1515 device: &B::Device,
1516 get_tensor: &F,
1517 prefix: &str,
1518 embedding_dim: usize,
1519 ) -> Result<Self>
1520 where
1521 F: Fn(&str) -> Option<(Vec<f32>, Vec<usize>)>,
1522 {
1523 let emb_prefix = format!("{}.emb", prefix);
1524 let emb = PixArtAlphaCombinedFlowEmbeddings::load_from_tensors(
1525 device,
1526 get_tensor,
1527 &emb_prefix,
1528 embedding_dim,
1529 )?;
1530
1531 let linear = load_linear_from_tensors(
1532 device,
1533 get_tensor,
1534 &format!("{}.linear", prefix),
1535 embedding_dim,
1536 6 * embedding_dim,
1537 )?;
1538
1539 Ok(Self { emb, linear })
1540 }
1541
1542 pub fn forward(&self, timestep: f32, hidden_dtype: DType) -> (Tensor<B, 3>, Tensor<B, 2>) {
1543 use burn::tensor::activation::silu;
1544
1545 let embedded_timestep = self.emb.forward(timestep, hidden_dtype);
1546 let timestep_mod = self.linear.forward(silu(embedded_timestep.clone()));
1547 let [batch, features] = timestep_mod.dims();
1548 (
1549 timestep_mod.reshape([batch, 6, features / 6]),
1550 embedded_timestep,
1551 )
1552 }
1553}
1554
1555#[derive(Module, Debug)]
1556pub struct PixArtAlphaCombinedFlowEmbeddings<B: Backend> {
1557 pub timestep_embedder: TimestepEmbedding<B>,
1558}
1559
1560impl<B: Backend> PixArtAlphaCombinedFlowEmbeddings<B> {
1561 pub fn new(device: &B::Device, embedding_dim: usize) -> Self {
1562 Self {
1563 timestep_embedder: TimestepEmbedding::new(device, 512, embedding_dim),
1564 }
1565 }
1566
1567 pub fn load_from_tensors<F>(
1568 device: &B::Device,
1569 get_tensor: &F,
1570 prefix: &str,
1571 embedding_dim: usize,
1572 ) -> Result<Self>
1573 where
1574 F: Fn(&str) -> Option<(Vec<f32>, Vec<usize>)>,
1575 {
1576 let timestep_embedder = TimestepEmbedding::load_from_tensors(
1577 device,
1578 get_tensor,
1579 &format!("{}.timestep_embedder", prefix),
1580 512,
1581 embedding_dim,
1582 )?;
1583
1584 Ok(Self { timestep_embedder })
1585 }
1586
1587 pub fn forward(&self, timestep: f32, _hidden_dtype: DType) -> Tensor<B, 2> {
1588 let device = self.timestep_embedder.linear_1.weight.device();
1589 let flow_t_size = 512;
1590 let half = flow_t_size / 2;
1591 let mut data = Vec::with_capacity(flow_t_size);
1592 for i in 0..half {
1593 let freq = 10_000_f32.powf(-(i as f32) / half.max(1) as f32);
1594 let angle = timestep * freq * 1000.0;
1595 data.push(angle.cos());
1596 }
1597 for i in 0..half {
1598 let freq = 10_000_f32.powf(-(i as f32) / half.max(1) as f32);
1599 let angle = timestep * freq * 1000.0;
1600 data.push(angle.sin());
1601 }
1602 let timestep = Tensor::<B, 2>::from_data(TensorData::new(data, [1, flow_t_size]), &device);
1603 self.timestep_embedder.forward(timestep)
1604 }
1605}
1606
1607#[derive(Module, Debug)]
1608pub struct TimestepEmbedding<B: Backend> {
1609 pub linear_1: Linear<B>,
1610 pub linear_2: Linear<B>,
1611}
1612
1613impl<B: Backend> TimestepEmbedding<B> {
1614 pub fn new(device: &B::Device, in_channels: usize, time_embed_dim: usize) -> Self {
1615 Self {
1616 linear_1: LinearConfig::new(in_channels, time_embed_dim)
1617 .with_bias(true)
1618 .with_layout(LinearLayout::Col)
1619 .init(device),
1620 linear_2: LinearConfig::new(time_embed_dim, time_embed_dim)
1621 .with_bias(true)
1622 .with_layout(LinearLayout::Col)
1623 .init(device),
1624 }
1625 }
1626
1627 pub fn load_from_tensors<F>(
1628 device: &B::Device,
1629 get_tensor: &F,
1630 prefix: &str,
1631 in_channels: usize,
1632 time_embed_dim: usize,
1633 ) -> Result<Self>
1634 where
1635 F: Fn(&str) -> Option<(Vec<f32>, Vec<usize>)>,
1636 {
1637 let linear_1 = load_linear_from_tensors(
1638 device,
1639 get_tensor,
1640 &format!("{}.linear_1", prefix),
1641 in_channels,
1642 time_embed_dim,
1643 )?;
1644
1645 let linear_2 = load_linear_from_tensors(
1646 device,
1647 get_tensor,
1648 &format!("{}.linear_2", prefix),
1649 time_embed_dim,
1650 time_embed_dim,
1651 )?;
1652
1653 Ok(Self { linear_1, linear_2 })
1654 }
1655
1656 pub fn forward(&self, x: Tensor<B, 2>) -> Tensor<B, 2> {
1657 use burn::tensor::activation::silu;
1658
1659 let emb = self.linear_1.forward(x);
1660 self.linear_2.forward(silu(emb))
1661 }
1662}
1663
1664#[derive(Module, Debug)]
1665pub struct ProjectLayer<B: Backend> {
1666 pub ffn_1: Conv1d<B>,
1667 pub ffn_2: Linear<B>,
1668 pub kernel_size: usize,
1669 pub out_channels: usize,
1670}
1671
1672impl<B: Backend> ProjectLayer<B> {
1673 pub fn new(
1674 device: &B::Device,
1675 in_channels: usize,
1676 out_channels: usize,
1677 kernel_size: usize,
1678 ) -> Self {
1679 let padding = kernel_size / 2;
1680 Self {
1681 ffn_1: Conv1dConfig::new(in_channels, out_channels, kernel_size)
1682 .with_padding(PaddingConfig1d::Explicit(padding, padding))
1683 .init(device),
1684 ffn_2: LinearConfig::new(out_channels, out_channels)
1685 .with_bias(true)
1686 .with_layout(LinearLayout::Col)
1687 .init(device),
1688 kernel_size,
1689 out_channels,
1690 }
1691 }
1692
1693 pub fn forward(&self, x: Tensor<B, 3>, _step: usize) -> Tensor<B, 3> {
1694 let x_t = x.swap_dims(1, 2);
1695
1696 let conv_out = self.forward_conv1d_exact(x_t);
1697
1698 let conv_out_t = conv_out.swap_dims(1, 2) * (self.kernel_size as f32).powf(-0.5);
1699
1700 self.ffn_2.forward(conv_out_t)
1701 }
1702
1703 fn forward_conv1d_exact(&self, input: Tensor<B, 3>) -> Tensor<B, 3> {
1704 let [batch, in_channels, seq_len] = input.dims();
1705 let [out_channels, weight_in_channels, kernel_size] = self.ffn_1.weight.dims();
1706 assert_eq!(
1707 in_channels, weight_in_channels,
1708 "project conv input channels mismatch"
1709 );
1710
1711 let padding = kernel_size / 2;
1712 let device = input.device();
1713 let left = Tensor::<B, 3>::zeros([batch, in_channels, padding], &device);
1714 let right = Tensor::<B, 3>::zeros([batch, in_channels, padding], &device);
1715 let padded = Tensor::cat(vec![left, input, right], 2);
1716
1717 let mut out = Tensor::<B, 3>::zeros([batch, seq_len, out_channels], &device);
1718 for k in 0..kernel_size {
1719 let x_k = padded
1720 .clone()
1721 .slice([0..batch, 0..in_channels, k..k + seq_len])
1722 .swap_dims(1, 2);
1723 let w_k = self
1724 .ffn_1
1725 .weight
1726 .val()
1727 .slice([0..out_channels, 0..in_channels, k..k + 1])
1728 .reshape([out_channels, in_channels])
1729 .swap_dims(0, 1);
1730 let x_k_flat = x_k.reshape([batch * seq_len, in_channels]);
1731 let projected = x_k_flat.matmul(w_k).reshape([batch, seq_len, out_channels]);
1732 out = out + projected;
1733 }
1734
1735 if let Some(bias) = &self.ffn_1.bias {
1736 let bias =
1737 bias.val()
1738 .reshape([1, 1, out_channels])
1739 .expand([batch, seq_len, out_channels]);
1740 out = out + bias;
1741 }
1742
1743 out.swap_dims(1, 2)
1744 }
1745
1746 pub fn load_from_tensors<F>(
1747 device: &B::Device,
1748 get_tensor: &F,
1749 prefix: &str,
1750 in_channels: usize,
1751 out_channels: usize,
1752 ) -> Result<Self>
1753 where
1754 F: Fn(&str) -> Option<(Vec<f32>, Vec<usize>)>,
1755 {
1756 use burn::nn::conv::Conv1dConfig;
1757
1758 let kernel_size = 3;
1759 let padding = kernel_size / 2;
1760
1761 let ffn_1_weight_name = format!("{}.ffn_1.weight", prefix);
1762 let ffn_1_bias_name = format!("{}.ffn_1.bias", prefix);
1763
1764 let ffn_1 = if let (Some((w_data, w_shape)), Some((b_data, b_shape))) =
1765 (get_tensor(&ffn_1_weight_name), get_tensor(&ffn_1_bias_name))
1766 {
1767 if w_shape.len() == 3
1768 && w_shape[0] == out_channels
1769 && w_shape[1] == in_channels
1770 && w_shape[2] == kernel_size
1771 && b_shape.len() == 1
1772 && b_shape[0] == out_channels
1773 {
1774 let weight = Tensor::<B, 3>::from_data(
1775 TensorData::new(w_data, [out_channels, in_channels, kernel_size]),
1776 device,
1777 );
1778 let bias =
1779 Tensor::<B, 1>::from_data(TensorData::new(b_data, [out_channels]), device);
1780
1781 Conv1d {
1782 weight: Param::from_tensor(weight),
1783 bias: Some(Param::from_tensor(bias)),
1784 stride: 1,
1785 kernel_size,
1786 dilation: 1,
1787 groups: 1,
1788 padding: PaddingConfig1d::Explicit(padding, padding),
1789 }
1790 } else {
1791 Conv1dConfig::new(in_channels, out_channels, kernel_size)
1792 .with_padding(PaddingConfig1d::Explicit(padding, padding))
1793 .init(device)
1794 }
1795 } else {
1796 Conv1dConfig::new(in_channels, out_channels, kernel_size)
1797 .with_padding(PaddingConfig1d::Explicit(padding, padding))
1798 .init(device)
1799 };
1800
1801 let ffn_2 = load_linear_from_tensors(
1802 device,
1803 get_tensor,
1804 &format!("{}.ffn_2", prefix),
1805 out_channels,
1806 out_channels,
1807 )?;
1808
1809 Ok(Self {
1810 ffn_1,
1811 ffn_2,
1812 kernel_size,
1813 out_channels,
1814 })
1815 }
1816}
1817
1818#[derive(Module, Debug)]
1819pub struct ScalarModel<B: Backend> {
1820 pub decoder_0: WNConv1d<B>,
1821 pub decoder_1: ResDecoderBlock<B>,
1822 pub decoder_2: ResDecoderBlock<B>,
1823 pub decoder_3: ResDecoderBlock<B>,
1824 pub decoder_4: ResDecoderBlock<B>,
1825 pub decoder_5: ResDecoderBlock<B>,
1826 pub decoder_6: PostProcessor<B>,
1827 pub decoder_7: WNConv1d<B>,
1828}
1829
1830impl<B: Backend> ScalarModel<B> {
1831 pub fn from_burnpack(path: &std::path::Path, device: &B::Device) -> Result<Self> {
1832 let mut model = Self::new(device, &HeartCodecConfig::default());
1833 let mut store = BurnpackStore::from_file(path).zero_copy(true);
1834
1835 if model.load_from(&mut store).is_ok() {
1836 return Ok(model);
1837 }
1838
1839 let get_tensor = HeartCodecModel::<B>::snapshots_to_f32_lookup(path)?;
1840 Self::load_from_dot_notation(path, device, &get_tensor)
1841 }
1842
1843 pub fn new(device: &B::Device, _config: &HeartCodecConfig) -> Self {
1844 let config = HeartCodecConfig::default();
1845 Self {
1846 decoder_0: WNConv1d::new(
1847 device,
1848 128,
1849 2048,
1850 config.delay_kernel_size,
1851 1,
1852 config.delay_kernel_size / 2,
1853 1,
1854 1,
1855 false,
1856 ),
1857
1858 decoder_1: ResDecoderBlock::new(device, 2048, 1024),
1859 decoder_2: ResDecoderBlock::new(device, 1024, 512),
1860 decoder_3: ResDecoderBlock::new(device, 512, 256),
1861 decoder_4: ResDecoderBlock::new(device, 256, 128),
1862 decoder_5: ResDecoderBlock::new(device, 128, 64),
1863
1864 decoder_6: PostProcessor::new(device, 64, 2),
1865
1866 decoder_7: WNConv1d::new(device, 64, 1, 7, 1, 3, 1, 1, true),
1867 }
1868 }
1869
1870 pub fn load_from_dot_notation<F>(
1871 _path: &std::path::Path,
1872 device: &B::Device,
1873 get_tensor: &F,
1874 ) -> Result<Self>
1875 where
1876 F: Fn(&str) -> Option<(Vec<f32>, Vec<usize>)>,
1877 {
1878 use crate::heartcodec::conv::{WNConv1dLoadArgs, load_wnconv_from_tensors};
1879
1880 let load_conv = |prefix: &str,
1881 in_ch: usize,
1882 out_ch: usize,
1883 ksize: usize,
1884 padding: usize,
1885 causal: bool|
1886 -> WNConv1d<B> {
1887 let weight_g_name = format!("{}.weight_g", prefix);
1888 let weight_v_name = format!("{}.weight_v", prefix);
1889 let bias_name = format!("{}.bias", prefix);
1890
1891 if let (Some(weight_g), Some(weight_v)) =
1892 (get_tensor(&weight_g_name), get_tensor(&weight_v_name))
1893 {
1894 let bias = get_tensor(&bias_name);
1895 return load_wnconv_from_tensors(
1896 device,
1897 WNConv1dLoadArgs {
1898 in_channels: in_ch,
1899 out_channels: out_ch,
1900 kernel_size: ksize,
1901 dilation: 1,
1902 causal,
1903 weight_g,
1904 weight_v,
1905 bias,
1906 },
1907 )
1908 .unwrap_or_else(|_| {
1909 WNConv1d::new(device, in_ch, out_ch, ksize, 1, padding, 1, 1, causal)
1910 });
1911 }
1912
1913 let g_name = format!("{}.parametrizations.weight.original0", prefix);
1914 let v_name = format!("{}.parametrizations.weight.original1", prefix);
1915
1916 if let (Some(weight_g), Some(weight_v)) = (get_tensor(&g_name), get_tensor(&v_name)) {
1917 let bias = get_tensor(&bias_name);
1918 return load_wnconv_from_tensors(
1919 device,
1920 WNConv1dLoadArgs {
1921 in_channels: in_ch,
1922 out_channels: out_ch,
1923 kernel_size: ksize,
1924 dilation: 1,
1925 causal,
1926 weight_g,
1927 weight_v,
1928 bias,
1929 },
1930 )
1931 .unwrap_or_else(|_| {
1932 WNConv1d::new(device, in_ch, out_ch, ksize, 1, padding, 1, 1, causal)
1933 });
1934 }
1935
1936 WNConv1d::new(device, in_ch, out_ch, ksize, 1, padding, 1, 1, causal)
1937 };
1938
1939 let decoder_0 = load_conv("scalar_model.decoder.0", 128, 2048, 5, 2, false);
1940 let decoder_1 = ResDecoderBlock::load_from_dot_notation(
1941 device,
1942 get_tensor,
1943 "scalar_model.decoder.1",
1944 2048,
1945 1024,
1946 )?;
1947 let decoder_2 = ResDecoderBlock::load_from_dot_notation(
1948 device,
1949 get_tensor,
1950 "scalar_model.decoder.2",
1951 1024,
1952 512,
1953 )?;
1954 let decoder_3 = ResDecoderBlock::load_from_dot_notation(
1955 device,
1956 get_tensor,
1957 "scalar_model.decoder.3",
1958 512,
1959 256,
1960 )?;
1961 let decoder_4 = ResDecoderBlock::load_from_dot_notation(
1962 device,
1963 get_tensor,
1964 "scalar_model.decoder.4",
1965 256,
1966 128,
1967 )?;
1968 let decoder_5 = ResDecoderBlock::load_from_dot_notation(
1969 device,
1970 get_tensor,
1971 "scalar_model.decoder.5",
1972 128,
1973 64,
1974 )?;
1975 let decoder_6 = if let Ok(pp) =
1976 PostProcessor::load_from_tensors(device, get_tensor, "scalar_model.decoder.6", 64, 2)
1977 {
1978 pp
1979 } else {
1980 PostProcessor::new(device, 64, 2)
1981 };
1982
1983 let decoder_7 = load_conv("scalar_model.decoder.7", 64, 1, 7, 3, true);
1984
1985 Ok(Self {
1986 decoder_0,
1987 decoder_1,
1988 decoder_2,
1989 decoder_3,
1990 decoder_4,
1991 decoder_5,
1992 decoder_6,
1993 decoder_7,
1994 })
1995 }
1996
1997 pub fn decode(&self, x: Tensor<B, 3>) -> Tensor<B, 3> {
1998 let x_quantized = (x.clone() * 9.0).round() / 9.0;
1999
2000 let h = self.decoder_0.forward(x_quantized);
2001 let h = self.decoder_1.forward(h);
2002 let h = self.decoder_2.forward(h);
2003 let h = self.decoder_3.forward(h);
2004 let h = self.decoder_4.forward(h);
2005 let h = self.decoder_5.forward(h);
2006 let h = self.decoder_6.forward(h);
2007 self.decoder_7.forward(h)
2008 }
2009
2010 pub fn decode_with_sync(&self, x: Tensor<B, 3>) -> Tensor<B, 3> {
2011 let x_quantized = (x.clone() * 9.0).round() / 9.0;
2012
2013 let h = self.decoder_0.forward(x_quantized);
2014 let _ = h.to_data();
2015
2016 let h = self.decoder_1.forward(h);
2017 let _ = h.to_data();
2018
2019 let h = self.decoder_2.forward(h);
2020 let _ = h.to_data();
2021
2022 let h = self.decoder_3.forward(h);
2023 let _ = h.to_data();
2024
2025 let h = self.decoder_4.forward(h);
2026 let _ = h.to_data();
2027
2028 let h = self.decoder_5.forward(h);
2029 let _ = h.to_data();
2030
2031 let h = self.decoder_6.forward(h);
2032 let _ = h.to_data();
2033
2034 self.decoder_7.forward(h)
2035 }
2036
2037 pub fn decode_latent_with_sync(&self, latent: Tensor<B, 3>) -> Tensor<B, 3> {
2038 let [batch, seq_len, channels] = latent.dims();
2039 assert_eq!(
2040 channels, 256,
2041 "Expected 256 channels from flow matching latent"
2042 );
2043
2044 let latent_reshaped = latent.reshape([batch, seq_len, 2, 128]);
2045 let latent_permuted = latent_reshaped.swap_dims(1, 2);
2046 let latent_split = latent_permuted.reshape([batch * 2, seq_len, 128]);
2047
2048 let scalar_input = latent_split.swap_dims(1, 2);
2049 self.decode_with_sync(scalar_input)
2050 }
2051}
2052
2053#[derive(Module, Debug)]
2054pub struct ResDecoderBlock<B: Backend> {
2055 pub up_conv: WNConvTranspose1d<B>,
2056 pub convs: Vec<ResidualUnit<B>>,
2057}
2058
2059impl<B: Backend> ResDecoderBlock<B> {
2060 pub fn new(device: &B::Device, in_ch: usize, out_ch: usize) -> Self {
2061 let (kernel_size, stride) = match (in_ch, out_ch) {
2062 (2048, 1024) => (10, 5),
2063 (1024, 512) => (8, 4),
2064 (512, 256) => (8, 4),
2065 (256, 128) => (8, 4),
2066 (128, 64) => (6, 3),
2067 _ => (8, 4),
2068 };
2069
2070 let up_conv = WNConvTranspose1d::new(
2071 device,
2072 in_ch,
2073 out_ch,
2074 kernel_size,
2075 stride,
2076 kernel_size / 2,
2077 0,
2078 1,
2079 1,
2080 true,
2081 );
2082
2083 let dilations = [1, 3, 5, 7, 9];
2084 let convs: Vec<_> = dilations
2085 .into_iter()
2086 .map(|dilation| ResidualUnit::new(device, out_ch, dilation))
2087 .collect();
2088
2089 Self { up_conv, convs }
2090 }
2091
2092 pub fn load_from_dot_notation<F>(
2093 device: &B::Device,
2094 get_tensor: &F,
2095 prefix: &str,
2096 in_ch: usize,
2097 out_ch: usize,
2098 ) -> Result<Self>
2099 where
2100 F: Fn(&str) -> Option<(Vec<f32>, Vec<usize>)>,
2101 {
2102 use crate::heartcodec::conv::{
2103 WNConvTranspose1dLoadArgs, load_wnconv_transpose_from_tensors,
2104 };
2105
2106 let (kernel_size, stride) = match (in_ch, out_ch) {
2107 (2048, 1024) => (10, 5),
2108 (1024, 512) => (8, 4),
2109 (512, 256) => (8, 4),
2110 (256, 128) => (8, 4),
2111 (128, 64) => (6, 3),
2112 _ => (8, 4),
2113 };
2114
2115 let up_conv_prefix = format!("{}.up_conv", prefix);
2116
2117 let up_conv = {
2118 let weight_g_name = format!("{}.weight_g", up_conv_prefix);
2119 let weight_v_name = format!("{}.weight_v", up_conv_prefix);
2120 let bias_name = format!("{}.layer.bias", up_conv_prefix);
2121
2122 if let (Some(weight_g), Some(weight_v)) =
2123 (get_tensor(&weight_g_name), get_tensor(&weight_v_name))
2124 {
2125 let bias = get_tensor(&bias_name);
2126 load_wnconv_transpose_from_tensors(
2127 device,
2128 WNConvTranspose1dLoadArgs {
2129 out_channels: out_ch,
2130 kernel_size,
2131 stride,
2132 causal: true,
2133 weight_g,
2134 weight_v,
2135 bias,
2136 },
2137 )
2138 .unwrap_or_else(|_| {
2139 WNConvTranspose1d::new(
2140 device,
2141 in_ch,
2142 out_ch,
2143 kernel_size,
2144 stride,
2145 kernel_size / 2,
2146 0,
2147 1,
2148 1,
2149 true,
2150 )
2151 })
2152 } else {
2153 let g_name = format!("{}.layer.parametrizations.weight.original0", up_conv_prefix);
2154 let v_name = format!("{}.layer.parametrizations.weight.original1", up_conv_prefix);
2155
2156 if let (Some(weight_g), Some(weight_v)) = (get_tensor(&g_name), get_tensor(&v_name))
2157 {
2158 let bias = get_tensor(&bias_name);
2159 load_wnconv_transpose_from_tensors(
2160 device,
2161 WNConvTranspose1dLoadArgs {
2162 out_channels: out_ch,
2163 kernel_size,
2164 stride,
2165 causal: true,
2166 weight_g,
2167 weight_v,
2168 bias,
2169 },
2170 )
2171 .unwrap_or_else(|_| {
2172 WNConvTranspose1d::new(
2173 device,
2174 in_ch,
2175 out_ch,
2176 kernel_size,
2177 stride,
2178 kernel_size / 2,
2179 0,
2180 1,
2181 1,
2182 true,
2183 )
2184 })
2185 } else {
2186 WNConvTranspose1d::new(
2187 device,
2188 in_ch,
2189 out_ch,
2190 kernel_size,
2191 stride,
2192 kernel_size / 2,
2193 0,
2194 1,
2195 1,
2196 true,
2197 )
2198 }
2199 }
2200 };
2201
2202 let mut convs = Vec::new();
2203 let dilations = [1, 3, 5, 7, 9];
2204 for (i, dilation) in dilations.into_iter().enumerate() {
2205 let unit_prefix = format!("{}.convs.{}", prefix, i);
2206 convs.push(ResidualUnit::load_from_dot_notation(
2207 device,
2208 get_tensor,
2209 &unit_prefix,
2210 out_ch,
2211 dilation,
2212 )?);
2213 }
2214
2215 Ok(Self { up_conv, convs })
2216 }
2217
2218 pub fn forward(&self, x: Tensor<B, 3>) -> Tensor<B, 3> {
2219 let mut h = self.up_conv.forward(x);
2220 for block in &self.convs {
2221 h = block.forward(h);
2222 }
2223 h
2224 }
2225}
2226
2227#[derive(Module, Debug)]
2228pub struct ResidualUnit<B: Backend> {
2229 pub conv1: WNConv1d<B>,
2230 pub conv2: WNConv1d<B>,
2231 pub activation1: PReLU<B>,
2232 pub activation2: PReLU<B>,
2233}
2234
2235impl<B: Backend> ResidualUnit<B> {
2236 pub fn new(device: &B::Device, channels: usize, dilation: usize) -> Self {
2237 Self {
2238 conv1: WNConv1d::new(
2239 device,
2240 channels,
2241 channels,
2242 7,
2243 1,
2244 3 * dilation,
2245 dilation,
2246 1,
2247 true,
2248 ),
2249 conv2: WNConv1d::new(device, channels, channels, 1, 1, 0, 1, 1, true),
2250 activation1: PReLU::new(device),
2251 activation2: PReLU::new(device),
2252 }
2253 }
2254
2255 pub fn load_from_dot_notation<F>(
2256 device: &B::Device,
2257 get_tensor: &F,
2258 prefix: &str,
2259 channels: usize,
2260 dilation: usize,
2261 ) -> Result<Self>
2262 where
2263 F: Fn(&str) -> Option<(Vec<f32>, Vec<usize>)>,
2264 {
2265 use crate::heartcodec::conv::{
2266 WNConv1dLoadArgs, load_prelu_from_tensor, load_wnconv_from_tensors,
2267 };
2268
2269 let load_conv = |conv_prefix: &str, ksize: usize, dilation: usize| -> WNConv1d<B> {
2270 let weight_g_name = format!("{}.weight_g", conv_prefix);
2271 let weight_v_name = format!("{}.weight_v", conv_prefix);
2272 let bias_name = format!("{}.bias", conv_prefix);
2273
2274 if let (Some(weight_g), Some(weight_v)) =
2275 (get_tensor(&weight_g_name), get_tensor(&weight_v_name))
2276 {
2277 let bias = get_tensor(&bias_name);
2278 return load_wnconv_from_tensors(
2279 device,
2280 WNConv1dLoadArgs {
2281 in_channels: channels,
2282 out_channels: channels,
2283 kernel_size: ksize,
2284 dilation,
2285 causal: true,
2286 weight_g,
2287 weight_v,
2288 bias,
2289 },
2290 )
2291 .unwrap_or_else(|_| {
2292 WNConv1d::new(
2293 device,
2294 channels,
2295 channels,
2296 ksize,
2297 1,
2298 (ksize / 2) * dilation,
2299 dilation,
2300 1,
2301 true,
2302 )
2303 });
2304 }
2305
2306 let g_name = format!("{}.parametrizations.weight.original0", conv_prefix);
2307 let v_name = format!("{}.parametrizations.weight.original1", conv_prefix);
2308 let bias_name = format!("{}.bias", conv_prefix);
2309
2310 if let (Some(weight_g), Some(weight_v)) = (get_tensor(&g_name), get_tensor(&v_name)) {
2311 let bias = get_tensor(&bias_name);
2312 return load_wnconv_from_tensors(
2313 device,
2314 WNConv1dLoadArgs {
2315 in_channels: channels,
2316 out_channels: channels,
2317 kernel_size: ksize,
2318 dilation,
2319 causal: true,
2320 weight_g,
2321 weight_v,
2322 bias,
2323 },
2324 )
2325 .unwrap_or_else(|_| {
2326 WNConv1d::new(
2327 device,
2328 channels,
2329 channels,
2330 ksize,
2331 1,
2332 (ksize / 2) * dilation,
2333 dilation,
2334 1,
2335 true,
2336 )
2337 });
2338 }
2339
2340 WNConv1d::new(
2341 device,
2342 channels,
2343 channels,
2344 ksize,
2345 1,
2346 (ksize / 2) * dilation,
2347 dilation,
2348 1,
2349 true,
2350 )
2351 };
2352
2353 let conv1_prefix = format!("{}.conv1", prefix);
2354 let conv2_prefix = format!("{}.conv2", prefix);
2355
2356 let conv1 = load_conv(&conv1_prefix, 7, dilation);
2357 let conv2 = load_conv(&conv2_prefix, 1, 1);
2358
2359 let act1_name = format!("{}.activation1.weight", prefix);
2360 let act2_name = format!("{}.activation2.weight", prefix);
2361
2362 let activation1 = if let Some((data, shape)) = get_tensor(&act1_name) {
2363 load_prelu_from_tensor(device, data, shape).unwrap_or_else(|_| PReLU::new(device))
2364 } else {
2365 PReLU::new(device)
2366 };
2367
2368 let activation2 = if let Some((data, shape)) = get_tensor(&act2_name) {
2369 load_prelu_from_tensor(device, data, shape).unwrap_or_else(|_| PReLU::new(device))
2370 } else {
2371 PReLU::new(device)
2372 };
2373
2374 Ok(Self {
2375 conv1,
2376 conv2,
2377 activation1,
2378 activation2,
2379 })
2380 }
2381
2382 pub fn forward(&self, x: Tensor<B, 3>) -> Tensor<B, 3> {
2383 let h = self.activation1.forward(self.conv1.forward(x.clone()));
2384 let h = self.activation2.forward(self.conv2.forward(h));
2385 x + h
2386 }
2387}
2388
2389#[derive(Module, Debug)]
2390pub struct PReLU<B: Backend> {
2391 pub weight: Param<Tensor<B, 1>>,
2392}
2393
2394impl<B: Backend> PReLU<B> {
2395 pub fn new(device: &B::Device) -> Self {
2396 Self {
2397 weight: Param::from_tensor(Tensor::ones([1], device) * 0.25),
2398 }
2399 }
2400
2401 pub fn forward(&self, x: Tensor<B, 3>) -> Tensor<B, 3> {
2402 use burn::tensor::activation::relu;
2403
2404 let weight = self.weight.val().reshape([1, 1, 1]);
2405 let positive = relu(x.clone());
2406 let negative = relu(x.neg()).neg() * weight;
2407 positive + negative
2408 }
2409}
2410
2411pub fn write_wav_from_f32(samples: &[f32], sample_rate: u32, path: &std::path::Path) -> Result<()> {
2412 write_wav_f32_oxideav(path, samples, 1, sample_rate)
2413 .with_context(|| format!("failed to write mono WAV to {}", path.display()))
2414}
2415
2416pub fn write_wav_from_f32_interleaved(
2417 samples: &[f32],
2418 channels: usize,
2419 frames: usize,
2420 sample_rate: u32,
2421 path: &std::path::Path,
2422) -> Result<()> {
2423 let sample_count = channels
2424 .checked_mul(frames)
2425 .context("interleaved float WAV sample count overflow")?;
2426 let mut interleaved = vec![0.0_f32; sample_count];
2427 interleaved
2428 .par_iter_mut()
2429 .enumerate()
2430 .for_each(|(sample_index, slot)| {
2431 let frame = sample_index / channels;
2432 let channel = sample_index % channels;
2433 let source_index = channel
2434 .checked_mul(frames)
2435 .and_then(|base| base.checked_add(frame))
2436 .unwrap_or(0);
2437 *slot = samples.get(source_index).copied().unwrap_or_default();
2438 });
2439 write_wav_f32_oxideav(path, &interleaved, channels, sample_rate)
2440 .with_context(|| format!("failed to write interleaved WAV to {}", path.display()))
2441}
2442
2443fn write_wav_f32_oxideav(
2444 path: &std::path::Path,
2445 samples: &[f32],
2446 channels: usize,
2447 sample_rate: u32,
2448) -> std::io::Result<()> {
2449 let channels = channels.max(1);
2450 if sample_rate == 0 {
2451 return Err(std::io::Error::other(
2452 "write_wav_f32: sample_rate must be > 0",
2453 ));
2454 }
2455 if channels > 8 {
2456 return Err(std::io::Error::other(format!(
2457 "write_wav_f32: channel count {channels} exceeds the supported maximum of 8"
2458 )));
2459 }
2460 if !samples.len().is_multiple_of(channels) {
2461 return Err(std::io::Error::other(
2462 "write_wav_f32: sample slice length is not a multiple of channels",
2463 ));
2464 }
2465
2466 let mut ctx = RuntimeContext::new();
2467 oxideav_basic::register(&mut ctx);
2468
2469 let stream = wav_f32_stream_info(channels, sample_rate);
2470 let file = std::fs::File::create(path)?;
2471 let output: Box<dyn oxideav_core::WriteSeek> = Box::new(file);
2472 let mut mux = ctx
2473 .containers
2474 .open_muxer("wav", output, std::slice::from_ref(&stream))
2475 .map_err(|e| std::io::Error::other(format!("OxideAV error: {e}")))?;
2476 mux.write_header()
2477 .map_err(|e| std::io::Error::other(format!("OxideAV error: {e}")))?;
2478
2479 let bytes: Vec<u8> = samples.iter().flat_map(|s| s.to_le_bytes()).collect();
2480 let packet = Packet::new(0, TimeBase::new(1, sample_rate as i64), bytes);
2481 mux.write_packet(&packet)
2482 .map_err(|e| std::io::Error::other(format!("OxideAV error: {e}")))?;
2483 mux.write_trailer()
2484 .map_err(|e| std::io::Error::other(format!("OxideAV error: {e}")))?;
2485 Ok(())
2486}
2487
2488fn wav_f32_stream_info(channels: usize, sample_rate: u32) -> StreamInfo {
2489 let mut params = CodecParameters::audio(CodecId::new("pcm_f32le"));
2490 params.media_type = MediaType::Audio;
2491 params.channels = Some(channels as u16);
2492 params.sample_rate = Some(sample_rate);
2493 params.sample_format = Some(SampleFormat::F32);
2494 StreamInfo {
2495 index: 0,
2496 time_base: TimeBase::new(1, sample_rate as i64),
2497 duration: None,
2498 start_time: Some(0),
2499 params,
2500 }
2501}
2502
2503pub fn frames_to_tensor<B: Backend>(frames: &[Vec<i64>], device: &B::Device) -> Tensor<B, 3, Int> {
2504 let num_frames = frames.len();
2505 let num_codebooks = if num_frames > 0 { frames[0].len() } else { 8 };
2506
2507 let mut data = Vec::with_capacity(num_frames * num_codebooks);
2508 for codebook in 0..num_codebooks {
2509 for frame in frames {
2510 if frame.len() != num_codebooks {
2511 panic!(
2512 "frames_to_tensor: inconsistent codebook count {}, expected {}",
2513 frame.len(),
2514 num_codebooks
2515 );
2516 }
2517 data.push(frame[codebook]);
2518 }
2519 }
2520
2521 Tensor::from_data(
2522 TensorData::new(data, [1, num_codebooks, num_frames]),
2523 device,
2524 )
2525}
2526
2527fn load_linear_from_tensors<B: Backend, F>(
2528 device: &B::Device,
2529 get_tensor: &F,
2530 prefix: &str,
2531 in_dim: usize,
2532 out_dim: usize,
2533) -> Result<Linear<B>>
2534where
2535 F: Fn(&str) -> Option<(Vec<f32>, Vec<usize>)>,
2536{
2537 let weight_name = format!("{}.weight", prefix);
2538 let bias_name = format!("{}.bias", prefix);
2539
2540 let weight = if let Some((data, shape)) = get_tensor(&weight_name) {
2541 if shape.len() == 2 && shape[0] == out_dim && shape[1] == in_dim {
2542 let mut transposed = vec![0.0f32; in_dim * out_dim];
2543 for i in 0..out_dim {
2544 for j in 0..in_dim {
2545 transposed[j * out_dim + i] = data[i * in_dim + j];
2546 }
2547 }
2548 Tensor::<B, 2>::from_data(TensorData::new(transposed, [in_dim, out_dim]), device)
2549 } else if shape.len() == 2 && shape[0] == in_dim && shape[1] == out_dim {
2550 Tensor::<B, 2>::from_data(TensorData::new(data, [in_dim, out_dim]), device)
2551 } else {
2552 Tensor::zeros([in_dim, out_dim], device)
2553 }
2554 } else {
2555 Tensor::zeros([in_dim, out_dim], device)
2556 };
2557
2558 let bias = if let Some((data, shape)) = get_tensor(&bias_name) {
2559 if shape.len() == 1 && shape[0] == out_dim {
2560 Some(Tensor::<B, 1>::from_data(
2561 TensorData::new(data, [out_dim]),
2562 device,
2563 ))
2564 } else {
2565 Some(Tensor::zeros([out_dim], device))
2566 }
2567 } else {
2568 None
2569 };
2570
2571 Ok(Linear {
2572 weight: Param::from_tensor(weight),
2573 bias: bias.map(Param::from_tensor),
2574 })
2575}
2576
2577fn load_rmsnorm_from_tensors<B: Backend, F>(
2578 device: &B::Device,
2579 get_tensor: &F,
2580 prefix: &str,
2581 dim: usize,
2582) -> Result<RmsNorm<B>>
2583where
2584 F: Fn(&str) -> Option<(Vec<f32>, Vec<usize>)>,
2585{
2586 let weight_name = format!("{}.weight", prefix);
2587
2588 let weight = if let Some((data, shape)) = get_tensor(&weight_name) {
2589 if shape.len() == 1 && shape[0] == dim {
2590 Tensor::<B, 1>::from_data(TensorData::new(data, [dim]), device)
2591 } else {
2592 Tensor::ones([dim], device)
2593 }
2594 } else {
2595 Tensor::ones([dim], device)
2596 };
2597
2598 Ok(RmsNorm {
2599 weight: Param::from_tensor(weight),
2600 })
2601}
2602
2603fn load_param_tensor<B: Backend, F, const D: usize>(
2604 device: &B::Device,
2605 get_tensor: &F,
2606 name: &str,
2607 shape: [usize; D],
2608) -> Result<Param<Tensor<B, D>>>
2609where
2610 F: Fn(&str) -> Option<(Vec<f32>, Vec<usize>)>,
2611{
2612 use burn::module::Param;
2613 use burn::tensor::TensorData;
2614
2615 let tensor = if let Some((data, _)) = get_tensor(name) {
2616 Tensor::from_data(TensorData::new(data, shape), device)
2617 } else {
2618 Tensor::zeros(shape, device)
2619 };
2620
2621 Ok(Param::from_tensor(tensor))
2622}
2623
2624#[cfg(test)]
2625mod tests {
2626 use super::{Tensor, TensorData, write_wav_from_f32, write_wav_from_f32_interleaved};
2627 use std::time::{SystemTime, UNIX_EPOCH};
2628
2629 fn temp_wav_path(name: &str) -> std::path::PathBuf {
2630 let nanos = SystemTime::now()
2631 .duration_since(UNIX_EPOCH)
2632 .expect("clock should be after epoch")
2633 .as_nanos();
2634 std::env::temp_dir().join(format!("maolan_{name}_{nanos}.wav"))
2635 }
2636
2637 #[derive(Debug, Clone)]
2638 struct WavInfo {
2639 channels: u16,
2640 sample_rate: u32,
2641 bits_per_sample: u16,
2642 is_float: bool,
2643 samples: Vec<f32>,
2644 }
2645
2646 fn read_wav_f32(path: &std::path::Path) -> WavInfo {
2647 use symphonia::core::audio::SampleBuffer;
2648 use symphonia::core::codecs::{CODEC_TYPE_NULL, DecoderOptions};
2649 use symphonia::core::errors::Error as SymphoniaError;
2650 use symphonia::core::formats::FormatOptions;
2651 use symphonia::core::io::MediaSourceStream;
2652 use symphonia::core::meta::MetadataOptions;
2653 use symphonia::core::probe::Hint;
2654
2655 let file = std::fs::File::open(path).expect("open wav");
2656 let mss = MediaSourceStream::new(Box::new(file), Default::default());
2657 let mut hint = Hint::new();
2658 if let Some(ext) = path.extension().and_then(|e| e.to_str()) {
2659 hint.with_extension(ext);
2660 }
2661 let probed = symphonia::default::get_probe()
2662 .format(
2663 &hint,
2664 mss,
2665 &FormatOptions::default(),
2666 &MetadataOptions::default(),
2667 )
2668 .expect("probe wav");
2669 let mut format = probed.format;
2670 let track = format
2671 .tracks()
2672 .iter()
2673 .find(|t| t.codec_params.codec != CODEC_TYPE_NULL)
2674 .or_else(|| format.tracks().first())
2675 .expect("audio track");
2676 let channels = track.codec_params.channels.map(|c| c.count()).unwrap_or(1);
2677 let sample_rate = track.codec_params.sample_rate.unwrap_or(48_000);
2678 let bits_per_sample = track.codec_params.bits_per_sample.unwrap_or(32);
2679 let sample_format = track.codec_params.sample_format;
2680 let is_float = sample_format
2681 .map(|f| {
2682 matches!(
2683 f,
2684 symphonia::core::sample::SampleFormat::F32
2685 | symphonia::core::sample::SampleFormat::F64
2686 )
2687 })
2688 .unwrap_or(true);
2689 let track_id = track.id;
2690 let mut decoder = symphonia::default::get_codecs()
2691 .make(&track.codec_params, &DecoderOptions::default())
2692 .expect("create decoder");
2693 let mut sample_buf = None;
2694 let mut samples = Vec::new();
2695 loop {
2696 let packet = match format.next_packet() {
2697 Ok(packet) => packet,
2698 Err(SymphoniaError::IoError(e))
2699 if e.kind() == std::io::ErrorKind::UnexpectedEof =>
2700 {
2701 break;
2702 }
2703 Err(e) => panic!("read error: {e}"),
2704 };
2705 if packet.track_id() != track_id {
2706 continue;
2707 }
2708 let decoded = decoder.decode(&packet).expect("decode packet");
2709 if sample_buf.is_none() {
2710 let spec = *decoded.spec();
2711 sample_buf = Some(SampleBuffer::<f32>::new(decoded.capacity() as u64, spec));
2712 }
2713 let buf = sample_buf.as_mut().unwrap();
2714 buf.copy_interleaved_ref(decoded);
2715 samples.extend_from_slice(buf.samples());
2716 }
2717 WavInfo {
2718 channels: channels as u16,
2719 sample_rate,
2720 bits_per_sample: bits_per_sample as u16,
2721 is_float,
2722 samples,
2723 }
2724 }
2725
2726 #[test]
2727 fn writes_mono_wav_as_f32() {
2728 let path = temp_wav_path("mono_f32");
2729 write_wav_from_f32(&[0.25, -0.5, 1.25], 48_000, &path).expect("write mono wav");
2730
2731 let info = read_wav_f32(&path);
2732 assert_eq!(info.channels, 1);
2733 assert_eq!(info.sample_rate, 48_000);
2734 assert_eq!(info.bits_per_sample, 32);
2735 assert!(info.is_float);
2736 assert_eq!(info.samples, vec![0.25, -0.5, 1.25]);
2737
2738 std::fs::remove_file(&path).expect("remove mono wav");
2739 }
2740
2741 #[test]
2742 fn writes_interleaved_wav_as_f32() {
2743 let path = temp_wav_path("stereo_f32");
2744 let planar_samples = [0.1, 0.2, 0.3, -0.1, -0.2, -0.3];
2745 write_wav_from_f32_interleaved(&planar_samples, 2, 3, 48_000, &path)
2746 .expect("write stereo wav");
2747
2748 let info = read_wav_f32(&path);
2749 assert_eq!(info.channels, 2);
2750 assert_eq!(info.sample_rate, 48_000);
2751 assert_eq!(info.bits_per_sample, 32);
2752 assert!(info.is_float);
2753 assert_eq!(info.samples, vec![0.1, -0.1, 0.2, -0.2, 0.3, -0.3]);
2754
2755 std::fs::remove_file(&path).expect("remove stereo wav");
2756 }
2757
2758 #[test]
2759 fn heartcodec_config_default_values() {
2760 let config = super::HeartCodecConfig::default();
2761 assert_eq!(config.dim, 512);
2762 assert_eq!(config.codebook_size, 8192);
2763 assert_eq!(config.codebook_dim, 32);
2764 assert_eq!(config.num_quantizers, 8);
2765 assert_eq!(config.attention_head_dim, 64);
2766 assert_eq!(config.in_channels, 1024);
2767 assert_eq!(config.num_attention_heads, 24);
2768 assert_eq!(config.num_layers, 24);
2769 assert_eq!(config.num_layers_2, 6);
2770 assert_eq!(config.out_channels, 256);
2771 assert_eq!(config.sample_rate, 48000);
2772 assert_eq!(config.latent_hidden_dim, 128);
2773 assert_eq!(config.init_channel, 64);
2774 assert_eq!(config.num_bands, 1);
2775 assert_eq!(config.num_samples, 2);
2776 assert_eq!(config.downsample_factors, [3, 4, 4, 4, 5]);
2777 assert_eq!(config.downsample_kernel_sizes, [6, 8, 8, 8, 10]);
2778 assert_eq!(config.upsample_factors, [5, 4, 4, 4, 3]);
2779 assert_eq!(config.upsample_kernel_sizes, [10, 8, 8, 8, 6]);
2780 assert_eq!(config.default_kernel_size, 7);
2781 assert_eq!(config.delay_kernel_size, 5);
2782 assert_eq!(config.res_kernel_size, 7);
2783 assert!(config.causal);
2784 assert_eq!(config.ode_steps, 10);
2785 }
2786
2787 #[test]
2788 fn frames_to_tensor_converts_correctly() {
2789 use burn::backend::ndarray::NdArray;
2790 let device = burn::prelude::Device::<NdArray<f32>>::default();
2791 let frames: Vec<Vec<i64>> = vec![
2792 vec![1, 2, 3, 4, 5, 6, 7, 8],
2793 vec![9, 10, 11, 12, 13, 14, 15, 16],
2794 vec![17, 18, 19, 20, 21, 22, 23, 24],
2795 ];
2796
2797 let tensor = super::frames_to_tensor::<NdArray<f32>>(&frames, &device);
2798 let dims = tensor.dims();
2799 assert_eq!(dims[0], 1);
2800 assert_eq!(dims[1], 8);
2801 assert_eq!(dims[2], 3);
2802 }
2803
2804 #[test]
2805 fn frames_to_tensor_empty_frames() {
2806 use burn::backend::ndarray::NdArray;
2807
2808 let device = burn::prelude::Device::<NdArray<f32>>::default();
2809 let frames: Vec<Vec<i64>> = vec![];
2810
2811 let tensor = super::frames_to_tensor::<NdArray<f32>>(&frames, &device);
2812 let dims = tensor.dims();
2813 assert_eq!(dims[0], 1);
2814 assert_eq!(dims[1], 8);
2815 assert_eq!(dims[2], 0);
2816 }
2817
2818 #[test]
2819 #[should_panic(expected = "inconsistent codebook count")]
2820 fn frames_to_tensor_rejects_inconsistent_codebooks() {
2821 use burn::backend::ndarray::NdArray;
2822
2823 let device = burn::prelude::Device::<NdArray<f32>>::default();
2824 let frames: Vec<Vec<i64>> = vec![vec![1, 2, 3, 4, 5, 6, 7, 8], vec![9, 10, 11]];
2825
2826 let _tensor = super::frames_to_tensor::<NdArray<f32>>(&frames, &device);
2827 }
2828
2829 #[test]
2830 fn prelu_activation_forward() {
2831 use burn::backend::ndarray::NdArray;
2832
2833 let device = burn::prelude::Device::<NdArray<f32>>::default();
2834 let prelu = super::PReLU::<NdArray<f32>>::new(&device);
2835
2836 let input = Tensor::<NdArray<f32>, 3>::from_data(
2837 TensorData::new(vec![1.0, 2.0, 3.0], [1, 1, 3]),
2838 &device,
2839 );
2840 let output = prelu.forward(input);
2841 let data = output.to_data().to_vec::<f32>().unwrap();
2842 assert!((data[0] - 1.0).abs() < 1e-6);
2843 assert!((data[1] - 2.0).abs() < 1e-6);
2844 assert!((data[2] - 3.0).abs() < 1e-6);
2845 }
2846
2847 #[test]
2848 fn prelu_activation_negative_values() {
2849 use burn::backend::ndarray::NdArray;
2850
2851 let device = burn::prelude::Device::<NdArray<f32>>::default();
2852 let prelu = super::PReLU::<NdArray<f32>>::new(&device);
2853
2854 let input = Tensor::<NdArray<f32>, 3>::from_data(
2855 TensorData::new(vec![-4.0, -8.0], [1, 1, 2]),
2856 &device,
2857 );
2858 let output = prelu.forward(input);
2859 let data = output.to_data().to_vec::<f32>().unwrap();
2860
2861 assert!((data[0] - (-1.0)).abs() < 1e-6);
2862 assert!((data[1] - (-2.0)).abs() < 1e-6);
2863 }
2864
2865 #[test]
2866 fn mlp_compute_hidden_dim() {
2867 let dim_1536 = super::Mlp::<burn::backend::ndarray::NdArray<f32>>::compute_hidden_dim(1536);
2868 let dim_3072 = super::Mlp::<burn::backend::ndarray::NdArray<f32>>::compute_hidden_dim(3072);
2869
2870 assert_eq!(dim_1536, 4096);
2871
2872 assert_eq!(dim_3072, 8192);
2873 }
2874
2875 #[test]
2876 fn rms_norm_forward_preserves_shape() {
2877 use burn::backend::ndarray::NdArray;
2878
2879 let device = burn::prelude::Device::<NdArray<f32>>::default();
2880 let rms_norm = super::RmsNorm::<NdArray<f32>>::new(&device, 16, 1e-6);
2881
2882 let input = Tensor::<NdArray<f32>, 3>::from_data(
2883 TensorData::new(vec![1.0; 128], [2, 4, 16]),
2884 &device,
2885 );
2886 let output = rms_norm.forward(input);
2887 assert_eq!(output.dims(), [2, 4, 16]);
2888 }
2889
2890 #[test]
2891 fn heartcodec_model_new_creates_valid_model() {
2892 use burn::backend::ndarray::NdArray;
2893
2894 let device = burn::prelude::Device::<NdArray<f32>>::default();
2895 let config = super::HeartCodecConfig::default();
2896 let model = super::HeartCodecModel::<NdArray<f32>>::new(&device);
2897
2898 assert_eq!(model.ode_steps, config.ode_steps);
2899 assert_eq!(model.guidance_scale, 1.0);
2900 }
2901
2902 #[test]
2903 fn heartcodec_model_with_ode_steps() {
2904 use burn::backend::ndarray::NdArray;
2905
2906 let device = burn::prelude::Device::<NdArray<f32>>::default();
2907 let model = super::HeartCodecModel::<NdArray<f32>>::new(&device).with_ode_steps(20);
2908
2909 assert_eq!(model.ode_steps, 20);
2910 }
2911
2912 #[test]
2913 fn heartcodec_model_with_guidance_scale() {
2914 use burn::backend::ndarray::NdArray;
2915
2916 let device = burn::prelude::Device::<NdArray<f32>>::default();
2917 let model = super::HeartCodecModel::<NdArray<f32>>::new(&device).with_guidance_scale(2.5);
2918
2919 assert_eq!(model.guidance_scale, 2.5);
2920 }
2921
2922 #[test]
2923 fn heartcodec_model_clamps_ode_steps() {
2924 use burn::backend::ndarray::NdArray;
2925
2926 let device = burn::prelude::Device::<NdArray<f32>>::default();
2927
2928 let model_high = super::HeartCodecModel::<NdArray<f32>>::new(&device).with_ode_steps(100);
2929 assert_eq!(model_high.ode_steps, 50);
2930
2931 let model_low = super::HeartCodecModel::<NdArray<f32>>::new(&device).with_ode_steps(0);
2932 assert_eq!(model_low.ode_steps, 1);
2933 }
2934
2935 #[test]
2936 fn residual_unit_new() {
2937 use burn::backend::ndarray::NdArray;
2938
2939 let device = burn::prelude::Device::<NdArray<f32>>::default();
2940 let unit = super::ResidualUnit::<NdArray<f32>>::new(&device, 128, 3);
2941
2942 assert_eq!(unit.conv1.dilation, 3);
2943 assert_eq!(unit.conv2.dilation, 1);
2944 }
2945
2946 #[test]
2947 fn resdecoder_block_new() {
2948 use burn::backend::ndarray::NdArray;
2949
2950 let device = burn::prelude::Device::<NdArray<f32>>::default();
2951 let block = super::ResDecoderBlock::<NdArray<f32>>::new(&device, 2048, 1024);
2952
2953 assert_eq!(block.up_conv.stride, 5);
2954 assert_eq!(block.convs.len(), 5);
2955 }
2956
2957 #[test]
2958 fn resdecoder_block_kernel_stride_pairs() {
2959 use burn::backend::ndarray::NdArray;
2960
2961 let device = burn::prelude::Device::<NdArray<f32>>::default();
2962
2963 let test_cases = [
2964 ((2048, 1024), (10, 5)),
2965 ((1024, 512), (8, 4)),
2966 ((512, 256), (8, 4)),
2967 ((256, 128), (8, 4)),
2968 ((128, 64), (6, 3)),
2969 ((100, 50), (8, 4)),
2970 ];
2971
2972 for ((in_ch, out_ch), (_expected_k, expected_s)) in test_cases {
2973 let block = super::ResDecoderBlock::<NdArray<f32>>::new(&device, in_ch, out_ch);
2974 assert_eq!(
2975 block.up_conv.stride, expected_s,
2976 "stride mismatch for ({}, {})",
2977 in_ch, out_ch
2978 );
2979 }
2980 }
2981
2982 #[test]
2983 fn project_layer_new() {
2984 use burn::backend::ndarray::NdArray;
2985
2986 let device = burn::prelude::Device::<NdArray<f32>>::default();
2987 let proj = super::ProjectLayer::<NdArray<f32>>::new(&device, 512, 256, 3);
2988
2989 assert_eq!(proj.kernel_size, 3);
2990 assert_eq!(proj.out_channels, 256);
2991 }
2992
2993 #[test]
2994 fn scalar_model_new() {
2995 use burn::backend::ndarray::NdArray;
2996
2997 let device = burn::prelude::Device::<NdArray<f32>>::default();
2998 let config = super::HeartCodecConfig::default();
2999 let model = super::ScalarModel::<NdArray<f32>>::new(&device, &config);
3000
3001 assert_eq!(model.decoder_0.kernel_size(), 5);
3002 assert_eq!(model.decoder_7.kernel_size(), 7);
3003 }
3004
3005 #[test]
3006 fn interpolate_1d_scale_factor_1() {
3007 use burn::backend::ndarray::NdArray;
3008
3009 let device = burn::prelude::Device::<NdArray<f32>>::default();
3010 let input = Tensor::<NdArray<f32>, 3>::from_data(
3011 TensorData::new(vec![1.0, 2.0, 3.0, 4.0], [1, 2, 2]),
3012 &device,
3013 );
3014
3015 let output = super::FlowMatching::<NdArray<f32>>::interpolate_1d(&input, 1);
3016 assert_eq!(output.dims(), input.dims());
3017 }
3018
3019 #[test]
3020 fn interpolate_1d_scale_factor_2() {
3021 use burn::backend::ndarray::NdArray;
3022
3023 let device = burn::prelude::Device::<NdArray<f32>>::default();
3024 let input = Tensor::<NdArray<f32>, 3>::from_data(
3025 TensorData::new(vec![1.0, 2.0], [1, 1, 2]),
3026 &device,
3027 );
3028
3029 let output = super::FlowMatching::<NdArray<f32>>::interpolate_1d(&input, 2);
3030 assert_eq!(output.dims(), [1, 2, 2]);
3031 }
3032
3033 #[test]
3034 fn build_latent_masks_basic() {
3035 let masks = super::FlowMatching::<burn::backend::ndarray::NdArray<f32>>::build_latent_masks(
3036 10, 8, 5,
3037 );
3038
3039 assert_eq!(masks.len(), 10);
3040
3041 assert_eq!(masks[0], 1);
3042 assert_eq!(masks[4], 1);
3043
3044 assert_eq!(masks[5], 2);
3045 assert_eq!(masks[7], 2);
3046
3047 assert_eq!(masks[8], 0);
3048 assert_eq!(masks[9], 0);
3049 }
3050
3051 #[test]
3052 fn build_latent_masks_incontext_larger_than_seq() {
3053 let masks = super::FlowMatching::<burn::backend::ndarray::NdArray<f32>>::build_latent_masks(
3054 5, 10, 10,
3055 );
3056
3057 assert_eq!(masks.len(), 5);
3058
3059 for mask in &masks {
3060 assert_eq!(*mask, 1);
3061 }
3062 }
3063
3064 #[test]
3065 fn transformer_block_new() {
3066 use burn::backend::ndarray::NdArray;
3067
3068 let device = burn::prelude::Device::<NdArray<f32>>::default();
3069 let block = super::TransformerBlock::<NdArray<f32>>::new(&device, 512, 8, 64);
3070
3071 assert_eq!(block.attn.num_heads, 8);
3072 assert_eq!(block.attn.head_dim, 64);
3073 }
3074
3075 #[test]
3076 fn attention_new() {
3077 use burn::backend::ndarray::NdArray;
3078
3079 let device = burn::prelude::Device::<NdArray<f32>>::default();
3080 let attn = super::Attention::<NdArray<f32>>::new(&device, 512, 8, 64);
3081
3082 assert_eq!(attn.num_heads, 8);
3083 assert_eq!(attn.head_dim, 64);
3084 assert_eq!(attn.rope_dim, 64);
3085 }
3086
3087 #[test]
3088 fn llama_transformer_new() {
3089 use burn::backend::ndarray::NdArray;
3090
3091 let device = burn::prelude::Device::<NdArray<f32>>::default();
3092 let config = super::HeartCodecConfig::default();
3093 let transformer = super::LlamaTransformer::<NdArray<f32>>::new(&device, &config);
3094
3095 assert_eq!(transformer.transformer_blocks.len(), 24);
3096 assert_eq!(transformer.transformer_blocks_2.len(), 6);
3097 }
3098
3099 #[test]
3100 fn adalayer_norm_single_new() {
3101 use burn::backend::ndarray::NdArray;
3102
3103 let device = burn::prelude::Device::<NdArray<f32>>::default();
3104 let adaln = super::AdaLayerNormSingle::<NdArray<f32>>::new(&device, 512);
3105
3106 let (output, embedded) = adaln.forward(0.5, burn::tensor::DType::F32);
3107 assert_eq!(output.dims(), [1, 6, 512]);
3108 assert_eq!(embedded.dims(), [1, 512]);
3109 }
3110
3111 #[test]
3112 fn timestep_embedding_new() {
3113 use burn::backend::ndarray::NdArray;
3114
3115 let device = burn::prelude::Device::<NdArray<f32>>::default();
3116 let emb = super::TimestepEmbedding::<NdArray<f32>>::new(&device, 128, 512);
3117
3118 let input = Tensor::<NdArray<f32>, 2>::zeros([1, 128], &device);
3119 let output = emb.forward(input);
3120 assert_eq!(output.dims(), [1, 512]);
3121 }
3122
3123 #[test]
3124 fn pixart_alpha_combined_flow_embeddings() {
3125 use burn::backend::ndarray::NdArray;
3126
3127 let device = burn::prelude::Device::<NdArray<f32>>::default();
3128 let emb = super::PixArtAlphaCombinedFlowEmbeddings::<NdArray<f32>>::new(&device, 512);
3129
3130 let output = emb.forward(0.5, burn::tensor::DType::F32);
3131 assert_eq!(output.dims(), [1, 512]);
3132 }
3133
3134 #[test]
3135 fn vq_codebook_new() {
3136 use burn::backend::ndarray::NdArray;
3137
3138 let device = burn::prelude::Device::<NdArray<f32>>::default();
3139 let codebook = super::VQCodebook::<NdArray<f32>>::new(&device, 1024, 64);
3140
3141 assert_eq!(codebook._codebook.embed.val().dims()[0], 1);
3142 assert_eq!(codebook._codebook.embed.val().dims()[1], 1024);
3143 assert_eq!(codebook._codebook.embed.val().dims()[2], 64);
3144 }
3145
3146 #[test]
3147 fn residual_vq_new() {
3148 use burn::backend::ndarray::NdArray;
3149
3150 let device = burn::prelude::Device::<NdArray<f32>>::default();
3151 let config = super::HeartCodecConfig::default();
3152 let rvq = super::ResidualVQ::<NdArray<f32>>::new(&device, &config);
3153
3154 assert_eq!(rvq.layers.len(), 8);
3155 }
3156}