Skip to main content

maolan_generate/heartcodec/
mod.rs

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}