Skip to main content

maolan_generate/
heartmula_runtime.rs

1use crate::heartcodec::frames_to_tensor;
2use anyhow::{Context, Result, anyhow};
3use burn::module::{
4    AutodiffModule, Content, Devices, EmptyRecord, Module, ModuleDisplay, ModuleDisplayDefault,
5    ModuleMapper, ModuleVisitor, Param, ParamId,
6};
7use burn::nn::{Embedding, EmbeddingConfig, Linear, LinearConfig, LinearLayout};
8use burn::prelude::Backend;
9use burn::tensor::activation::{silu, softmax};
10use burn::tensor::backend::AutodiffBackend;
11use burn::tensor::{Bool, DType, Int, Tensor, TensorData};
12use burn_store::{BurnpackStore, ModuleSnapshot, ModuleStore};
13use rayon::prelude::*;
14use serde::{Deserialize, Serialize};
15use std::env;
16use std::fs;
17use std::fs::File;
18use std::io::{BufReader, BufWriter, Read, Write};
19use std::path::{Path, PathBuf};
20use tokie::Tokenizer;
21
22const HEARTMULA_PARALLEL_TOKENS: usize = 9;
23const HEARTMULA_AUDIO_CODEBOOKS: usize = 8;
24const HEARTMULA_HIDDEN_SIZE: usize = 3072;
25const HEARTMULA_MUQ_DIM: usize = 512;
26const HEARTMULA_BACKBONE_LAYERS: usize = 28;
27const HEARTMULA_BACKBONE_HEADS: usize = 24;
28const HEARTMULA_BACKBONE_KV_HEADS: usize = 8;
29const HEARTMULA_DECODER_LAYERS: usize = 3;
30const HEARTMULA_DECODER_HEADS: usize = 8;
31const HEARTMULA_DECODER_KV_HEADS: usize = 4;
32const HEARTMULA_MLP_DIM: usize = 8192;
33const HEARTMULA_NORM_EPSILON: f64 = 1e-5;
34const HEARTMULA_ROPE_BASE: f32 = 500_000.0;
35const HEARTMULA_ROPE_SCALE_FACTOR: f32 = 32.0;
36const HEARTMULA_OLD_CONTEXT_LEN: f32 = 8192.0;
37const HEARTMULA_LOW_FREQ_FACTOR: f32 = 1.0;
38const HEARTMULA_HIGH_FREQ_FACTOR: f32 = 4.0;
39const HEARTCODEC_STAGE_ENV: &str = "MAOLAN_HEARTCODEC_STAGE";
40const HEARTCODEC_STAGE_FLOW: &str = "flow";
41const HEARTCODEC_STAGE_SCALAR: &str = "scalar";
42const HEARTCODEC_STAGE_PLAN_JSON_ENV: &str = "MAOLAN_HEARTCODEC_STAGE_PLAN_JSON";
43const HEARTCODEC_STAGE_PLAN_MAGIC: &[u8; 8] = b"MHCPLAN1";
44const HEARTCODEC_SEGMENT_DURATION_SECONDS: f32 = 29.76;
45
46type ProgressCallback<'a> = dyn FnMut(&str, f32, &str) + 'a;
47
48#[derive(Debug, Serialize, Deserialize)]
49pub struct HeartmulaJsonOutput {
50    pub model: String,
51    pub runtime: String,
52    pub tags: String,
53    pub lyrics: String,
54    pub frames: Vec<Vec<i64>>,
55    pub frame_count: usize,
56    pub sample_rate_hz: u32,
57}
58
59#[derive(Debug, Serialize, Deserialize)]
60pub struct HeartmulaFirstFrameDebug {
61    pub history: Vec<[i64; HEARTMULA_PARALLEL_TOKENS]>,
62    pub backbone_prefill_input_dims: Vec<usize>,
63    pub backbone_prefill_input: Vec<f32>,
64    pub backbone_layer0_prefill_hidden_dims: Vec<usize>,
65    pub backbone_layer0_prefill_hidden: Vec<f32>,
66    pub last_hidden_dims: Vec<usize>,
67    pub last_hidden: Vec<f32>,
68    pub backbone_layer0_prefill_q_dims: Vec<usize>,
69    pub backbone_layer0_prefill_q: Vec<f32>,
70    pub backbone_layer0_prefill_k_expanded_dims: Vec<usize>,
71    pub backbone_layer0_prefill_k_expanded: Vec<f32>,
72    pub backbone_layer0_prefill_v_expanded_dims: Vec<usize>,
73    pub backbone_layer0_prefill_v_expanded: Vec<f32>,
74    pub backbone_layer0_prefill_k_dims: Vec<usize>,
75    pub backbone_layer0_prefill_k: Vec<f32>,
76    pub backbone_layer0_prefill_v_dims: Vec<usize>,
77    pub backbone_layer0_prefill_v: Vec<f32>,
78    pub backbone_last_prefill_k_dims: Vec<usize>,
79    pub backbone_last_prefill_k: Vec<f32>,
80    pub backbone_last_prefill_v_dims: Vec<usize>,
81    pub backbone_last_prefill_v: Vec<f32>,
82    pub guided_codebook0_logits_dims: Vec<usize>,
83    pub guided_codebook0_logits: Vec<f32>,
84    pub argmax_first_frame: Vec<i64>,
85    pub second_history_row: Vec<i64>,
86    pub second_hidden_input_dims: Vec<usize>,
87    pub second_hidden_input: Vec<f32>,
88    pub second_layer0_q_dims: Vec<usize>,
89    pub second_layer0_q: Vec<f32>,
90    pub second_layer0_k_expanded_dims: Vec<usize>,
91    pub second_layer0_k_expanded: Vec<f32>,
92    pub second_layer0_v_expanded_dims: Vec<usize>,
93    pub second_layer0_v_expanded: Vec<f32>,
94    pub second_layer0_full_k_dims: Vec<usize>,
95    pub second_layer0_full_k: Vec<f32>,
96    pub second_layer0_full_v_dims: Vec<usize>,
97    pub second_layer0_full_v: Vec<f32>,
98    pub second_layer0_attn_out_dims: Vec<usize>,
99    pub second_layer0_attn_out: Vec<f32>,
100    pub second_layer0_mlp_out_dims: Vec<usize>,
101    pub second_layer0_mlp_out: Vec<f32>,
102    pub second_hidden_dims: Vec<usize>,
103    pub second_hidden: Vec<f32>,
104    pub second_layer_outputs_dims: Vec<Vec<usize>>,
105    pub second_layer_outputs: Vec<Vec<f32>>,
106    pub second_guided_codebook0_logits_dims: Vec<usize>,
107    pub second_guided_codebook0_logits: Vec<f32>,
108    pub second_argmax_frame: Vec<i64>,
109    pub second_decoder_step_inputs_dims: Vec<Vec<usize>>,
110    pub second_decoder_step_inputs: Vec<Vec<f32>>,
111    pub second_decoder_step_hidden_dims: Vec<Vec<usize>>,
112    pub second_decoder_step_hidden: Vec<Vec<f32>>,
113    pub second_guided_decoder_logits_dims: Vec<Vec<usize>>,
114    pub second_guided_decoder_logits: Vec<Vec<f32>>,
115    pub second_decoder_layer0_step2_q_dims: Vec<usize>,
116    pub second_decoder_layer0_step2_q: Vec<f32>,
117    pub second_decoder_layer0_step2_k_expanded_dims: Vec<usize>,
118    pub second_decoder_layer0_step2_k_expanded: Vec<f32>,
119    pub second_decoder_layer0_step2_v_expanded_dims: Vec<usize>,
120    pub second_decoder_layer0_step2_v_expanded: Vec<f32>,
121    pub second_decoder_layer0_step2_full_k_dims: Vec<usize>,
122    pub second_decoder_layer0_step2_full_k: Vec<f32>,
123    pub second_decoder_layer0_step2_full_v_dims: Vec<usize>,
124    pub second_decoder_layer0_step2_full_v: Vec<f32>,
125    pub guided_decoder_logits_dims: Vec<Vec<usize>>,
126    pub guided_decoder_logits: Vec<Vec<f32>>,
127    pub decoder_step_inputs_dims: Vec<Vec<usize>>,
128    pub decoder_step_inputs: Vec<Vec<f32>>,
129    pub decoder_step_hidden_dims: Vec<Vec<usize>>,
130    pub decoder_step_hidden: Vec<Vec<f32>>,
131}
132
133#[derive(Debug, Serialize, Deserialize)]
134struct LatentTensorFile {
135    dims: [usize; 3],
136    data: Vec<f32>,
137}
138
139pub struct HeartmulaGenerationConfig<'a> {
140    pub text_bos_id: i64,
141    pub text_eos_id: i64,
142    pub audio_eos_id: i64,
143    pub empty_id: i64,
144    pub lyrics_ids: &'a [i64],
145    pub tags_ids: &'a [i64],
146    pub max_audio_frames: usize,
147
148    pub temperature: f32,
149
150    pub topk: usize,
151
152    pub cfg_scale: f32,
153
154    pub progress_callback: Option<Box<ProgressCallback<'a>>>,
155}
156
157#[derive(Clone)]
158struct HeartmulaTransformerCache<B: Backend> {
159    layers: Vec<HeartmulaAttentionCache<B>>,
160}
161
162#[derive(Clone)]
163struct HeartmulaAttentionCache<B: Backend> {
164    key: Option<Tensor<B, 4>>,
165    value: Option<Tensor<B, 4>>,
166}
167
168#[derive(Clone, Debug)]
169struct SplitAudioEmbeddings<B: Backend> {
170    table: Option<Tensor<B, 2>>,
171    vocab_size: usize,
172}
173
174#[derive(Module, Debug)]
175pub struct HeartmulaModel<B: Backend> {
176    pub text_embeddings: Embedding<B>,
177    audio_embeddings: SplitAudioEmbeddings<B>,
178    pub unconditional_text_embedding: Embedding<B>,
179    pub projection: Linear<B>,
180    pub codebook0_head: Linear<B>,
181    pub audio_head: Param<Tensor<B, 3>>,
182    pub muq_linear: Linear<B>,
183    pub backbone: HeartmulaTransformer<B>,
184    pub decoder: HeartmulaTransformer<B>,
185}
186
187#[derive(Module, Debug)]
188pub struct HeartmulaTransformer<B: Backend> {
189    pub layers: Vec<HeartmulaTransformerLayer<B>>,
190    pub norm: HeartmulaRmsNorm<B>,
191}
192
193#[derive(Module, Debug)]
194pub struct HeartmulaTransformerLayer<B: Backend> {
195    pub attn: HeartmulaAttention<B>,
196    pub mlp: HeartmulaMlp<B>,
197    pub sa_norm: HeartmulaRmsNorm<B>,
198    pub mlp_norm: HeartmulaRmsNorm<B>,
199}
200
201#[derive(Module, Debug)]
202pub struct HeartmulaAttention<B: Backend> {
203    pub q_proj: Linear<B>,
204    pub k_proj: Linear<B>,
205    pub v_proj: Linear<B>,
206    pub output_proj: Linear<B>,
207    #[module(skip)]
208    meta: AttentionMeta,
209}
210
211#[derive(Module, Debug)]
212pub struct HeartmulaMlp<B: Backend> {
213    pub w1: Linear<B>,
214    pub w2: Linear<B>,
215    pub w3: Linear<B>,
216}
217
218#[derive(Module, Debug)]
219pub struct HeartmulaRmsNorm<B: Backend> {
220    pub scale: Param<Tensor<B, 1>>,
221    pub epsilon: f64,
222}
223
224#[derive(Clone, Debug)]
225struct AttentionMeta {
226    num_heads: usize,
227    num_kv_heads: usize,
228    head_dim: usize,
229}
230
231impl<B: Backend> HeartmulaModel<B> {
232    pub fn new(device: &B::Device, text_vocab_size: usize, audio_vocab_size: usize) -> Self {
233        Self {
234            text_embeddings: EmbeddingConfig::new(text_vocab_size, HEARTMULA_HIDDEN_SIZE)
235                .init(device),
236            audio_embeddings: SplitAudioEmbeddings::new_placeholder(audio_vocab_size),
237            unconditional_text_embedding: EmbeddingConfig::new(1, HEARTMULA_HIDDEN_SIZE)
238                .init(device),
239            projection: linear_no_bias(device, HEARTMULA_HIDDEN_SIZE, HEARTMULA_HIDDEN_SIZE),
240            codebook0_head: linear_no_bias(device, HEARTMULA_HIDDEN_SIZE, audio_vocab_size),
241            audio_head: uninitialized_param(
242                [
243                    HEARTMULA_AUDIO_CODEBOOKS - 1,
244                    HEARTMULA_HIDDEN_SIZE,
245                    audio_vocab_size,
246                ],
247                device,
248            ),
249            muq_linear: linear_with_bias(device, HEARTMULA_MUQ_DIM, HEARTMULA_HIDDEN_SIZE),
250            backbone: HeartmulaTransformer::new(
251                device,
252                HEARTMULA_BACKBONE_LAYERS,
253                HEARTMULA_BACKBONE_HEADS,
254                HEARTMULA_BACKBONE_KV_HEADS,
255            ),
256            decoder: HeartmulaTransformer::new(
257                device,
258                HEARTMULA_DECODER_LAYERS,
259                HEARTMULA_DECODER_HEADS,
260                HEARTMULA_DECODER_KV_HEADS,
261            ),
262        }
263    }
264
265    pub fn from_burnpack(
266        path: &Path,
267        device: &B::Device,
268        text_vocab_size: usize,
269        audio_vocab_size: usize,
270    ) -> Result<Self> {
271        let mut model = Self::new(device, text_vocab_size, audio_vocab_size);
272        let snapshots = BurnpackStore::from_file(path)
273            .zero_copy(true)
274            .get_all_snapshots()
275            .with_context(|| format!("failed to read snapshots from {}", path.display()))?
276            .clone();
277        let audio_embedding_data = snapshots
278            .iter()
279            .find_map(|(_, snap)| {
280                (snap.full_path() == "audio_embeddings.weight").then(|| snap.to_data().ok())
281            })
282            .flatten()
283            .ok_or_else(|| anyhow!("missing audio_embeddings.weight in HeartMula burnpack"))?
284            .convert::<f32>();
285        let mut store = BurnpackStore::from_file(path).zero_copy(true);
286        model
287            .load_from(&mut store)
288            .with_context(|| format!("failed to load HeartMula weights from {}", path.display()))?;
289        model.audio_embeddings =
290            SplitAudioEmbeddings::load_from_data(device, audio_embedding_data, audio_vocab_size)?;
291        Ok(model)
292    }
293
294    pub fn generate_frames(
295        &self,
296        device: &B::Device,
297        config: &mut HeartmulaGenerationConfig<'_>,
298    ) -> Result<Vec<Vec<i64>>> {
299        let normalized_tags =
300            normalize_text_ids(config.text_bos_id, config.text_eos_id, config.tags_ids);
301        let history = build_prompt_history(
302            config.text_bos_id,
303            config.text_eos_id,
304            config.lyrics_ids,
305            config.tags_ids,
306        );
307        let muq_index = normalized_tags.len();
308        let mut frames = Vec::new();
309        let mut backbone_cache = self.backbone.new_cache();
310        let mut last_hidden = self.prefill_backbone(
311            device,
312            &history,
313            Some(muq_index),
314            config.cfg_scale > 1.0,
315            &mut backbone_cache,
316        )?;
317        sync_and_cleanup_backend::<B>(device)?;
318
319        const CHUNK_SIZE: usize = 12;
320        let total_chunks = config.max_audio_frames.div_ceil(CHUNK_SIZE);
321
322        for chunk_idx in 0..total_chunks {
323            let chunk_start = chunk_idx * CHUNK_SIZE;
324            let chunk_end = ((chunk_idx + 1) * CHUNK_SIZE).min(config.max_audio_frames);
325            let frames_in_chunk = chunk_end - chunk_start;
326
327            let progress = (chunk_idx as f32 / total_chunks as f32) * 0.99;
328            if let Some(ref mut cb) = config.progress_callback {
329                cb("generator", progress, "Generating audio tokens");
330            }
331
332            for _ in 0..frames_in_chunk {
333                if frames.len() >= config.max_audio_frames {
334                    break;
335                }
336
337                let _frame_index = frames.len();
338                let next_frame = self.decode_frame_from_last_hidden(
339                    device,
340                    last_hidden.clone(),
341                    config.temperature,
342                    config.topk,
343                    config.cfg_scale,
344                )?;
345                if next_frame.iter().any(|token| *token >= config.audio_eos_id) {
346                    return Ok(frames);
347                }
348
349                frames.push(next_frame);
350                let next_row = build_audio_history_row(
351                    frames.last().expect("frame was just pushed"),
352                    config.empty_id,
353                );
354                let next_hidden = if config.cfg_scale > 1.0 {
355                    let hidden = self.embed_single_history_row(device, &next_row);
356                    Tensor::cat(vec![hidden.clone(), hidden], 0)
357                } else {
358                    self.embed_single_history_row(device, &next_row)
359                };
360                let next_position = (history.len() + frames.len() - 1) as i64;
361                last_hidden = self.backbone.forward_incremental(
362                    next_hidden,
363                    single_position_tensor::<B>(next_position, device),
364                    &mut backbone_cache,
365                )?;
366                sync_and_cleanup_backend::<B>(device)?;
367            }
368
369            sync_and_cleanup_backend::<B>(device)?;
370            std::thread::sleep(std::time::Duration::from_millis(50));
371        }
372
373        Ok(frames)
374    }
375
376    pub fn debug_first_frame(
377        &self,
378        device: &B::Device,
379        config: &mut HeartmulaGenerationConfig<'_>,
380    ) -> Result<HeartmulaFirstFrameDebug> {
381        let normalized_tags =
382            normalize_text_ids(config.text_bos_id, config.text_eos_id, config.tags_ids);
383        let history = build_prompt_history(
384            config.text_bos_id,
385            config.text_eos_id,
386            config.lyrics_ids,
387            config.tags_ids,
388        );
389        let muq_index = normalized_tags.len();
390        let tokens = history_tokens_tensor::<B>(&history, device);
391        let tokens_mask = history_mask_tensor::<B>(&history, device);
392        let history_hidden_cond = self.embed_history(tokens.clone(), tokens_mask.clone(), false);
393        let mut history_hidden_for_debug = if config.cfg_scale > 1.0 {
394            let history_hidden_uncond = self.embed_history(tokens, tokens_mask, true);
395            Tensor::cat(vec![history_hidden_cond, history_hidden_uncond], 0)
396        } else {
397            history_hidden_cond
398        };
399        if Some(muq_index) == Some(muq_index) {
400            let muq_zero = Tensor::<B, 2>::zeros([1, HEARTMULA_MUQ_DIM], device);
401            let muq_hidden =
402                self.muq_linear
403                    .forward(muq_zero)
404                    .reshape([1, 1, HEARTMULA_HIDDEN_SIZE]);
405            history_hidden_for_debug = if config.cfg_scale > 1.0 {
406                let uncond_hidden = self
407                    .unconditional_text_embedding
408                    .forward(Tensor::<B, 2, Int>::zeros([1, 1], device))
409                    .reshape([1, 1, HEARTMULA_HIDDEN_SIZE]);
410                let replacement = Tensor::cat(vec![muq_hidden, uncond_hidden], 0);
411                splice_sequence_token(history_hidden_for_debug, replacement, muq_index)
412            } else {
413                splice_sequence_token(history_hidden_for_debug, muq_hidden, muq_index)
414            };
415        }
416        let positions = position_tensor::<B>((0..history.len() as i64).collect(), device);
417        let layer0 = self
418            .backbone
419            .layers
420            .first()
421            .ok_or_else(|| anyhow!("missing backbone layer 0"))?;
422        let layer0_hidden = layer0.sa_norm.forward(history_hidden_for_debug.clone());
423        let [batch, seq_len, _] = layer0_hidden.dims();
424        let mut layer0_q = layer0.attn.q_proj.forward(layer0_hidden.clone()).reshape([
425            batch,
426            seq_len,
427            layer0.attn.meta.num_heads,
428            layer0.attn.meta.head_dim,
429        ]);
430        let mut layer0_k = layer0.attn.k_proj.forward(layer0_hidden.clone()).reshape([
431            batch,
432            seq_len,
433            layer0.attn.meta.num_kv_heads,
434            layer0.attn.meta.head_dim,
435        ]);
436        let mut layer0_v = layer0.attn.v_proj.forward(layer0_hidden.clone()).reshape([
437            batch,
438            seq_len,
439            layer0.attn.meta.num_kv_heads,
440            layer0.attn.meta.head_dim,
441        ]);
442        layer0_q = apply_scaled_rope(layer0_q, &positions);
443        layer0_k = apply_scaled_rope(layer0_k, &positions);
444        if layer0.attn.meta.num_heads != layer0.attn.meta.num_kv_heads {
445            let repeats = layer0.attn.meta.num_heads / layer0.attn.meta.num_kv_heads;
446            layer0_k = repeat_kv_heads(layer0_k, repeats);
447            layer0_v = repeat_kv_heads(layer0_v, repeats);
448        }
449        let mut backbone_cache = self.backbone.new_cache();
450        let last_hidden = self.prefill_backbone(
451            device,
452            &history,
453            Some(muq_index),
454            config.cfg_scale > 1.0,
455            &mut backbone_cache,
456        )?;
457        let layer0_cache = backbone_cache
458            .layers
459            .first()
460            .ok_or_else(|| anyhow!("missing backbone layer 0 cache"))?;
461        let prefill_k = layer0_cache
462            .key
463            .clone()
464            .ok_or_else(|| anyhow!("missing backbone layer 0 key cache after prefill"))?;
465        let prefill_v = layer0_cache
466            .value
467            .clone()
468            .ok_or_else(|| anyhow!("missing backbone layer 0 value cache after prefill"))?;
469        let last_layer_cache = backbone_cache
470            .layers
471            .last()
472            .ok_or_else(|| anyhow!("missing backbone last layer cache"))?;
473        let last_prefill_k = last_layer_cache
474            .key
475            .clone()
476            .ok_or_else(|| anyhow!("missing backbone last layer key cache after prefill"))?;
477        let last_prefill_v = last_layer_cache
478            .value
479            .clone()
480            .ok_or_else(|| anyhow!("missing backbone last layer value cache after prefill"))?;
481
482        let use_cfg = config.cfg_scale > 1.0;
483        let codebook0_logits = self.codebook0_head.forward(last_hidden.clone());
484        let guided_codebook0_logits = if use_cfg {
485            let cond_logits = codebook0_logits
486                .clone()
487                .slice([0..1, 0..self.audio_vocab_size()]);
488            let uncond_logits = codebook0_logits
489                .clone()
490                .slice([1..2, 0..self.audio_vocab_size()]);
491            uncond_logits.clone() + (cond_logits - uncond_logits) * config.cfg_scale
492        } else {
493            codebook0_logits
494        };
495        let use_cfg = config.cfg_scale > 1.0;
496        let argmax_first_frame = self.decode_frame_from_last_hidden(
497            device,
498            last_hidden.clone(),
499            1.0,
500            1,
501            config.cfg_scale,
502        )?;
503        let next_row = build_audio_history_row(&argmax_first_frame, config.empty_id);
504        let next_hidden = if use_cfg {
505            let hidden = self.embed_single_history_row(device, &next_row);
506            Tensor::cat(vec![hidden.clone(), hidden], 0)
507        } else {
508            self.embed_single_history_row(device, &next_row)
509        };
510        let next_position = history.len() as i64;
511        let second_layer0 = self
512            .backbone
513            .layers
514            .first()
515            .ok_or_else(|| anyhow!("missing backbone layer 0"))?;
516        let second_layer0_hidden = second_layer0.sa_norm.forward(next_hidden.clone());
517        let [second_batch, second_seq_len, _] = second_layer0_hidden.dims();
518        let mut second_layer0_q = second_layer0
519            .attn
520            .q_proj
521            .forward(second_layer0_hidden.clone())
522            .reshape([
523                second_batch,
524                second_seq_len,
525                second_layer0.attn.meta.num_heads,
526                second_layer0.attn.meta.head_dim,
527            ]);
528        let mut second_layer0_k_unrepeated = second_layer0
529            .attn
530            .k_proj
531            .forward(second_layer0_hidden.clone())
532            .reshape([
533                second_batch,
534                second_seq_len,
535                second_layer0.attn.meta.num_kv_heads,
536                second_layer0.attn.meta.head_dim,
537            ]);
538        let second_layer0_v_unrepeated = second_layer0
539            .attn
540            .v_proj
541            .forward(second_layer0_hidden.clone())
542            .reshape([
543                second_batch,
544                second_seq_len,
545                second_layer0.attn.meta.num_kv_heads,
546                second_layer0.attn.meta.head_dim,
547            ]);
548        let second_position_tensor = single_position_tensor::<B>(next_position, device);
549        second_layer0_q = apply_scaled_rope(second_layer0_q, &second_position_tensor);
550        second_layer0_k_unrepeated =
551            apply_scaled_rope(second_layer0_k_unrepeated, &second_position_tensor);
552        let mut second_layer0_k = second_layer0_k_unrepeated.clone();
553        let mut second_layer0_v = second_layer0_v_unrepeated.clone();
554        if second_layer0.attn.meta.num_heads != second_layer0.attn.meta.num_kv_heads {
555            let repeats = second_layer0.attn.meta.num_heads / second_layer0.attn.meta.num_kv_heads;
556            second_layer0_k = repeat_kv_heads(second_layer0_k, repeats);
557            second_layer0_v = repeat_kv_heads(second_layer0_v, repeats);
558        }
559        let second_layer0_q_swapped = second_layer0_q.clone().swap_dims(1, 2);
560        let second_layer0_k_unrepeated = second_layer0_k_unrepeated.swap_dims(1, 2);
561        let second_layer0_v_unrepeated = second_layer0_v_unrepeated.swap_dims(1, 2);
562        let second_full_k_unrepeated = Tensor::cat(
563            vec![prefill_k.clone(), second_layer0_k_unrepeated.clone()],
564            2,
565        );
566        let second_full_v_unrepeated = Tensor::cat(
567            vec![prefill_v.clone(), second_layer0_v_unrepeated.clone()],
568            2,
569        );
570        let second_full_k = if second_layer0.attn.meta.num_heads
571            != second_layer0.attn.meta.num_kv_heads
572        {
573            let repeats = second_layer0.attn.meta.num_heads / second_layer0.attn.meta.num_kv_heads;
574            repeat_cached_kv_heads(second_full_k_unrepeated.clone(), repeats)
575        } else {
576            second_full_k_unrepeated.clone()
577        };
578        let second_full_v = if second_layer0.attn.meta.num_heads
579            != second_layer0.attn.meta.num_kv_heads
580        {
581            let repeats = second_layer0.attn.meta.num_heads / second_layer0.attn.meta.num_kv_heads;
582            repeat_cached_kv_heads(second_full_v_unrepeated.clone(), repeats)
583        } else {
584            second_full_v_unrepeated.clone()
585        };
586        let second_scores = second_layer0_q_swapped
587            .clone()
588            .matmul(second_full_k.clone().swap_dims(2, 3))
589            .mul_scalar(1.0 / (second_layer0.attn.meta.head_dim as f32).sqrt());
590        let second_weights = softmax(second_scores, 3);
591        let second_attn_out = second_layer0.attn.output_proj.forward(
592            second_weights
593                .matmul(second_full_v.clone())
594                .swap_dims(1, 2)
595                .reshape([second_batch, second_seq_len, HEARTMULA_HIDDEN_SIZE]),
596        );
597        let second_layer0_after_attn = next_hidden.clone() + second_attn_out.clone();
598        let second_layer0_mlp_out = second_layer0.mlp.forward(
599            second_layer0
600                .mlp_norm
601                .forward(second_layer0_after_attn.clone()),
602        );
603        let mut second_layer_outputs_dims = Vec::with_capacity(self.backbone.layers.len());
604        let mut second_layer_outputs = Vec::with_capacity(self.backbone.layers.len());
605        let mut second_hidden_seq = next_hidden.clone();
606        for (layer, layer_cache) in self
607            .backbone
608            .layers
609            .iter()
610            .zip(backbone_cache.layers.iter_mut())
611        {
612            second_hidden_seq = layer.forward_incremental(
613                second_hidden_seq,
614                second_position_tensor.clone(),
615                layer_cache,
616            )?;
617            second_layer_outputs_dims.push(second_hidden_seq.dims().to_vec());
618            second_layer_outputs.push(tensor_to_f32_vec(second_hidden_seq.clone())?);
619        }
620        let second_hidden = take_last_token(self.backbone.norm.forward(second_hidden_seq));
621        let second_codebook0_logits = self.codebook0_head.forward(second_hidden.clone());
622        let second_guided_codebook0_logits = if use_cfg {
623            let cond_logits = second_codebook0_logits
624                .clone()
625                .slice([0..1, 0..self.audio_vocab_size()]);
626            let uncond_logits = second_codebook0_logits
627                .clone()
628                .slice([1..2, 0..self.audio_vocab_size()]);
629            uncond_logits.clone() + (cond_logits - uncond_logits) * config.cfg_scale
630        } else {
631            second_codebook0_logits
632        };
633        let second_argmax_frame = self.decode_frame_from_last_hidden(
634            device,
635            second_hidden.clone(),
636            1.0,
637            1,
638            config.cfg_scale,
639        )?;
640        let mut second_guided_decoder_logits_dims =
641            Vec::with_capacity(HEARTMULA_AUDIO_CODEBOOKS - 1);
642        let mut second_guided_decoder_logits = Vec::with_capacity(HEARTMULA_AUDIO_CODEBOOKS - 1);
643        let mut second_decoder_step_inputs_dims = Vec::with_capacity(HEARTMULA_AUDIO_CODEBOOKS - 1);
644        let mut second_decoder_step_inputs = Vec::with_capacity(HEARTMULA_AUDIO_CODEBOOKS - 1);
645        let mut second_decoder_step_hidden_dims = Vec::with_capacity(HEARTMULA_AUDIO_CODEBOOKS - 1);
646        let mut second_decoder_step_hidden = Vec::with_capacity(HEARTMULA_AUDIO_CODEBOOKS - 1);
647        let mut second_decoder_cache = self.decoder.new_cache();
648        let second_c0_token = *second_argmax_frame
649            .first()
650            .ok_or_else(|| anyhow!("second argmax frame was empty"))?;
651        let second_c0_embed = self.embed_audio_token(device, 0, second_c0_token);
652        let second_c0_embed = if use_cfg {
653            Tensor::cat(vec![second_c0_embed.clone(), second_c0_embed], 0)
654        } else {
655            second_c0_embed
656        };
657        let second_decoder_input = Tensor::cat(
658            vec![
659                second_hidden.clone().unsqueeze_dim(1),
660                second_c0_embed.clone(),
661            ],
662            1,
663        );
664        let second_decoder_input = self.projection.forward(second_decoder_input);
665        let second_first_decoder_h = self.decoder.forward_prefill(
666            second_decoder_input.clone(),
667            position_tensor::<B>(vec![0, 1], device),
668            &mut second_decoder_cache,
669        )?;
670        let mut second_decoder_layer0_step2_q_dims = Vec::new();
671        let mut second_decoder_layer0_step2_q = Vec::new();
672        let mut second_decoder_layer0_step2_k_expanded_dims = Vec::new();
673        let mut second_decoder_layer0_step2_k_expanded = Vec::new();
674        let mut second_decoder_layer0_step2_v_expanded_dims = Vec::new();
675        let mut second_decoder_layer0_step2_v_expanded = Vec::new();
676        let mut second_decoder_layer0_step2_full_k_dims = Vec::new();
677        let mut second_decoder_layer0_step2_full_k = Vec::new();
678        let mut second_decoder_layer0_step2_full_v_dims = Vec::new();
679        let mut second_decoder_layer0_step2_full_v = Vec::new();
680        let mut second_current_embed: Option<Tensor<B, 3>> = None;
681        let mut second_next_decoder_pos = 2_i64;
682        for codebook in 1..HEARTMULA_AUDIO_CODEBOOKS {
683            if codebook == 2 {
684                let embed = second_current_embed.clone().ok_or_else(|| {
685                    anyhow!("missing second decoder embed for codebook {}", codebook)
686                })?;
687                let second_decoder_input = self.projection.forward(embed.clone());
688                let layer0 = self
689                    .decoder
690                    .layers
691                    .first()
692                    .ok_or_else(|| anyhow!("missing decoder layer 0"))?;
693                let layer0_hidden = layer0.sa_norm.forward(second_decoder_input.clone());
694                let [batch, seq_len, _] = layer0_hidden.dims();
695                let mut q = layer0.attn.q_proj.forward(layer0_hidden.clone()).reshape([
696                    batch,
697                    seq_len,
698                    layer0.attn.meta.num_heads,
699                    layer0.attn.meta.head_dim,
700                ]);
701                let mut k_unrepeated = layer0.attn.k_proj.forward(layer0_hidden.clone()).reshape([
702                    batch,
703                    seq_len,
704                    layer0.attn.meta.num_kv_heads,
705                    layer0.attn.meta.head_dim,
706                ]);
707                let v_unrepeated = layer0.attn.v_proj.forward(layer0_hidden.clone()).reshape([
708                    batch,
709                    seq_len,
710                    layer0.attn.meta.num_kv_heads,
711                    layer0.attn.meta.head_dim,
712                ]);
713                let pos = single_position_tensor::<B>(second_next_decoder_pos, device);
714                q = apply_scaled_rope(q, &pos);
715                k_unrepeated = apply_scaled_rope(k_unrepeated, &pos);
716                let mut k = k_unrepeated.clone();
717                let mut v = v_unrepeated.clone();
718                if layer0.attn.meta.num_heads != layer0.attn.meta.num_kv_heads {
719                    let repeats = layer0.attn.meta.num_heads / layer0.attn.meta.num_kv_heads;
720                    k = repeat_kv_heads(k, repeats);
721                    v = repeat_kv_heads(v, repeats);
722                }
723                let q_swapped = q.clone().swap_dims(1, 2);
724                let k_unrepeated = k_unrepeated.swap_dims(1, 2);
725                let v_unrepeated = v_unrepeated.swap_dims(1, 2);
726                let layer0_cache = second_decoder_cache
727                    .layers
728                    .first()
729                    .ok_or_else(|| anyhow!("missing decoder layer 0 cache"))?;
730                let prev_k = layer0_cache
731                    .key
732                    .clone()
733                    .ok_or_else(|| anyhow!("missing decoder layer 0 key cache"))?;
734                let prev_v = layer0_cache
735                    .value
736                    .clone()
737                    .ok_or_else(|| anyhow!("missing decoder layer 0 value cache"))?;
738                let full_k_unrepeated = Tensor::cat(vec![prev_k, k_unrepeated.clone()], 2);
739                let full_v_unrepeated = Tensor::cat(vec![prev_v, v_unrepeated.clone()], 2);
740                let full_k = if layer0.attn.meta.num_heads != layer0.attn.meta.num_kv_heads {
741                    let repeats = layer0.attn.meta.num_heads / layer0.attn.meta.num_kv_heads;
742                    repeat_cached_kv_heads(full_k_unrepeated, repeats)
743                } else {
744                    full_k_unrepeated
745                };
746                let full_v = if layer0.attn.meta.num_heads != layer0.attn.meta.num_kv_heads {
747                    let repeats = layer0.attn.meta.num_heads / layer0.attn.meta.num_kv_heads;
748                    repeat_cached_kv_heads(full_v_unrepeated, repeats)
749                } else {
750                    full_v_unrepeated
751                };
752                second_decoder_layer0_step2_q_dims = q_swapped.dims().to_vec();
753                second_decoder_layer0_step2_q = tensor_to_f32_vec(q_swapped)?;
754                second_decoder_layer0_step2_k_expanded_dims = k.dims().to_vec();
755                second_decoder_layer0_step2_k_expanded = tensor_to_f32_vec(k)?;
756                second_decoder_layer0_step2_v_expanded_dims = v.dims().to_vec();
757                second_decoder_layer0_step2_v_expanded = tensor_to_f32_vec(v)?;
758                second_decoder_layer0_step2_full_k_dims = full_k.dims().to_vec();
759                second_decoder_layer0_step2_full_k = tensor_to_f32_vec(full_k)?;
760                second_decoder_layer0_step2_full_v_dims = full_v.dims().to_vec();
761                second_decoder_layer0_step2_full_v = tensor_to_f32_vec(full_v)?;
762            }
763            let head = self
764                .audio_head
765                .val()
766                .slice([
767                    codebook - 1..codebook,
768                    0..HEARTMULA_HIDDEN_SIZE,
769                    0..self.audio_head.dims()[2],
770                ])
771                .reshape([HEARTMULA_HIDDEN_SIZE, self.audio_head.dims()[2]]);
772            let logits = if codebook == 1 {
773                second_decoder_step_inputs_dims.push(second_decoder_input.dims().to_vec());
774                second_decoder_step_inputs.push(tensor_to_f32_vec(second_decoder_input.clone())?);
775                second_decoder_step_hidden_dims.push(second_first_decoder_h.dims().to_vec());
776                second_decoder_step_hidden.push(tensor_to_f32_vec(second_first_decoder_h.clone())?);
777                second_first_decoder_h.clone().matmul(head.clone())
778            } else {
779                let embed = second_current_embed.clone().ok_or_else(|| {
780                    anyhow!("missing second decoder embed for codebook {}", codebook)
781                })?;
782                let second_decoder_input = self.projection.forward(embed);
783                second_decoder_step_inputs_dims.push(second_decoder_input.dims().to_vec());
784                second_decoder_step_inputs.push(tensor_to_f32_vec(second_decoder_input.clone())?);
785                let second_last_decoder_h = self.decoder.forward_incremental(
786                    second_decoder_input,
787                    single_position_tensor::<B>(second_next_decoder_pos, device),
788                    &mut second_decoder_cache,
789                )?;
790                second_decoder_step_hidden_dims.push(second_last_decoder_h.dims().to_vec());
791                second_decoder_step_hidden.push(tensor_to_f32_vec(second_last_decoder_h.clone())?);
792                second_next_decoder_pos += 1;
793                second_last_decoder_h.matmul(head.clone())
794            };
795            let guided_logits = if use_cfg {
796                let cond_logits = logits.clone().slice([0..1, 0..self.audio_vocab_size()]);
797                let uncond_logits = logits.slice([1..2, 0..self.audio_vocab_size()]);
798                uncond_logits.clone() + (cond_logits - uncond_logits) * config.cfg_scale
799            } else {
800                logits
801            };
802            second_guided_decoder_logits_dims.push(guided_logits.dims().to_vec());
803            second_guided_decoder_logits.push(tensor_to_f32_vec(guided_logits.clone())?);
804            let token = *second_argmax_frame
805                .get(codebook)
806                .ok_or_else(|| anyhow!("missing second argmax token for codebook {}", codebook))?;
807            second_current_embed = Some(self.embed_audio_token(device, codebook, token));
808            if use_cfg {
809                let embed = second_current_embed
810                    .clone()
811                    .ok_or_else(|| anyhow!("missing second decoder embed after sampling"))?;
812                second_current_embed = Some(Tensor::cat(vec![embed.clone(), embed], 0));
813            }
814        }
815        let mut guided_decoder_logits_dims = Vec::with_capacity(HEARTMULA_AUDIO_CODEBOOKS - 1);
816        let mut guided_decoder_logits = Vec::with_capacity(HEARTMULA_AUDIO_CODEBOOKS - 1);
817        let mut decoder_step_inputs_dims = Vec::with_capacity(HEARTMULA_AUDIO_CODEBOOKS - 1);
818        let mut decoder_step_inputs = Vec::with_capacity(HEARTMULA_AUDIO_CODEBOOKS - 1);
819        let mut decoder_step_hidden_dims = Vec::with_capacity(HEARTMULA_AUDIO_CODEBOOKS - 1);
820        let mut decoder_step_hidden = Vec::with_capacity(HEARTMULA_AUDIO_CODEBOOKS - 1);
821
822        let mut decoder_cache = self.decoder.new_cache();
823        let c0_token = *argmax_first_frame
824            .first()
825            .ok_or_else(|| anyhow!("argmax first frame was empty"))?;
826        let c0_embed = self.embed_audio_token(device, 0, c0_token);
827        let c0_embed = if use_cfg {
828            Tensor::cat(vec![c0_embed.clone(), c0_embed], 0)
829        } else {
830            c0_embed
831        };
832        let decoder_input = Tensor::cat(
833            vec![last_hidden.clone().unsqueeze_dim(1), c0_embed.clone()],
834            1,
835        );
836        let decoder_input = self.projection.forward(decoder_input);
837        let first_decoder_h = self.decoder.forward_prefill(
838            decoder_input.clone(),
839            position_tensor::<B>(vec![0, 1], device),
840            &mut decoder_cache,
841        )?;
842        let mut current_embed: Option<Tensor<B, 3>> = None;
843        let mut next_decoder_pos = 2_i64;
844        for codebook in 1..HEARTMULA_AUDIO_CODEBOOKS {
845            let head = self
846                .audio_head
847                .val()
848                .slice([
849                    codebook - 1..codebook,
850                    0..HEARTMULA_HIDDEN_SIZE,
851                    0..self.audio_head.dims()[2],
852                ])
853                .reshape([HEARTMULA_HIDDEN_SIZE, self.audio_head.dims()[2]]);
854            let logits = if codebook == 1 {
855                decoder_step_inputs_dims.push(decoder_input.dims().to_vec());
856                decoder_step_inputs.push(tensor_to_f32_vec(decoder_input.clone())?);
857                decoder_step_hidden_dims.push(first_decoder_h.dims().to_vec());
858                decoder_step_hidden.push(tensor_to_f32_vec(first_decoder_h.clone())?);
859                first_decoder_h.clone().matmul(head.clone())
860            } else {
861                let embed = current_embed.clone().ok_or_else(|| {
862                    anyhow!("missing debug decoder embed for codebook {}", codebook)
863                })?;
864                let decoder_input = self.projection.forward(embed);
865                decoder_step_inputs_dims.push(decoder_input.dims().to_vec());
866                decoder_step_inputs.push(tensor_to_f32_vec(decoder_input.clone())?);
867                let last_decoder_h = self.decoder.forward_incremental(
868                    decoder_input,
869                    single_position_tensor::<B>(next_decoder_pos, device),
870                    &mut decoder_cache,
871                )?;
872                decoder_step_hidden_dims.push(last_decoder_h.dims().to_vec());
873                decoder_step_hidden.push(tensor_to_f32_vec(last_decoder_h.clone())?);
874                next_decoder_pos += 1;
875                last_decoder_h.matmul(head.clone())
876            };
877            let guided_logits = if use_cfg {
878                let cond_logits = logits.clone().slice([0..1, 0..self.audio_vocab_size()]);
879                let uncond_logits = logits.slice([1..2, 0..self.audio_vocab_size()]);
880                uncond_logits.clone() + (cond_logits - uncond_logits) * config.cfg_scale
881            } else {
882                logits
883            };
884            guided_decoder_logits_dims.push(guided_logits.dims().to_vec());
885            guided_decoder_logits.push(tensor_to_f32_vec(guided_logits.clone())?);
886            let token = *argmax_first_frame
887                .get(codebook)
888                .ok_or_else(|| anyhow!("missing argmax token for codebook {}", codebook))?;
889            current_embed = Some(self.embed_audio_token(device, codebook, token));
890            if use_cfg {
891                let embed = current_embed
892                    .clone()
893                    .ok_or_else(|| anyhow!("missing debug decoder embed after sampling"))?;
894                current_embed = Some(Tensor::cat(vec![embed.clone(), embed], 0));
895            }
896        }
897
898        Ok(HeartmulaFirstFrameDebug {
899            history,
900            backbone_prefill_input_dims: history_hidden_for_debug.dims().to_vec(),
901            backbone_prefill_input: tensor_to_f32_vec(history_hidden_for_debug)?,
902            backbone_layer0_prefill_hidden_dims: layer0_hidden.dims().to_vec(),
903            backbone_layer0_prefill_hidden: tensor_to_f32_vec(layer0_hidden)?,
904            last_hidden_dims: last_hidden.dims().to_vec(),
905            last_hidden: tensor_to_f32_vec(last_hidden)?,
906            backbone_layer0_prefill_q_dims: layer0_q.dims().to_vec(),
907            backbone_layer0_prefill_q: tensor_to_f32_vec(layer0_q)?,
908            backbone_layer0_prefill_k_expanded_dims: layer0_k.dims().to_vec(),
909            backbone_layer0_prefill_k_expanded: tensor_to_f32_vec(layer0_k.clone())?,
910            backbone_layer0_prefill_v_expanded_dims: layer0_v.dims().to_vec(),
911            backbone_layer0_prefill_v_expanded: tensor_to_f32_vec(layer0_v.clone())?,
912            backbone_layer0_prefill_k_dims: prefill_k.dims().to_vec(),
913            backbone_layer0_prefill_k: tensor_to_f32_vec(prefill_k)?,
914            backbone_layer0_prefill_v_dims: prefill_v.dims().to_vec(),
915            backbone_layer0_prefill_v: tensor_to_f32_vec(prefill_v)?,
916            backbone_last_prefill_k_dims: last_prefill_k.dims().to_vec(),
917            backbone_last_prefill_k: tensor_to_f32_vec(last_prefill_k)?,
918            backbone_last_prefill_v_dims: last_prefill_v.dims().to_vec(),
919            backbone_last_prefill_v: tensor_to_f32_vec(last_prefill_v)?,
920            guided_codebook0_logits_dims: guided_codebook0_logits.dims().to_vec(),
921            guided_codebook0_logits: tensor_to_f32_vec(guided_codebook0_logits)?,
922            argmax_first_frame,
923            second_history_row: next_row.to_vec(),
924            second_hidden_input_dims: next_hidden.dims().to_vec(),
925            second_hidden_input: tensor_to_f32_vec(next_hidden.clone())?,
926            second_layer0_q_dims: second_layer0_q_swapped.dims().to_vec(),
927            second_layer0_q: tensor_to_f32_vec(second_layer0_q_swapped)?,
928            second_layer0_k_expanded_dims: second_layer0_k.dims().to_vec(),
929            second_layer0_k_expanded: tensor_to_f32_vec(second_layer0_k)?,
930            second_layer0_v_expanded_dims: second_layer0_v.dims().to_vec(),
931            second_layer0_v_expanded: tensor_to_f32_vec(second_layer0_v)?,
932            second_layer0_full_k_dims: second_full_k.dims().to_vec(),
933            second_layer0_full_k: tensor_to_f32_vec(second_full_k)?,
934            second_layer0_full_v_dims: second_full_v.dims().to_vec(),
935            second_layer0_full_v: tensor_to_f32_vec(second_full_v)?,
936            second_layer0_attn_out_dims: second_attn_out.dims().to_vec(),
937            second_layer0_attn_out: tensor_to_f32_vec(second_attn_out)?,
938            second_layer0_mlp_out_dims: second_layer0_mlp_out.dims().to_vec(),
939            second_layer0_mlp_out: tensor_to_f32_vec(second_layer0_mlp_out)?,
940            second_hidden_dims: second_hidden.dims().to_vec(),
941            second_hidden: tensor_to_f32_vec(second_hidden)?,
942            second_layer_outputs_dims,
943            second_layer_outputs,
944            second_guided_codebook0_logits_dims: second_guided_codebook0_logits.dims().to_vec(),
945            second_guided_codebook0_logits: tensor_to_f32_vec(second_guided_codebook0_logits)?,
946            second_argmax_frame,
947            second_decoder_step_inputs_dims,
948            second_decoder_step_inputs,
949            second_decoder_step_hidden_dims,
950            second_decoder_step_hidden,
951            second_guided_decoder_logits_dims,
952            second_guided_decoder_logits,
953            second_decoder_layer0_step2_q_dims,
954            second_decoder_layer0_step2_q,
955            second_decoder_layer0_step2_k_expanded_dims,
956            second_decoder_layer0_step2_k_expanded,
957            second_decoder_layer0_step2_v_expanded_dims,
958            second_decoder_layer0_step2_v_expanded,
959            second_decoder_layer0_step2_full_k_dims,
960            second_decoder_layer0_step2_full_k,
961            second_decoder_layer0_step2_full_v_dims,
962            second_decoder_layer0_step2_full_v,
963            guided_decoder_logits_dims,
964            guided_decoder_logits,
965            decoder_step_inputs_dims,
966            decoder_step_inputs,
967            decoder_step_hidden_dims,
968            decoder_step_hidden,
969        })
970    }
971
972    fn prefill_backbone(
973        &self,
974        device: &B::Device,
975        history: &[[i64; HEARTMULA_PARALLEL_TOKENS]],
976        muq_insert_index: Option<usize>,
977        use_cfg: bool,
978        cache: &mut HeartmulaTransformerCache<B>,
979    ) -> Result<Tensor<B, 2>> {
980        let tokens = history_tokens_tensor::<B>(history, device);
981        let tokens_mask = history_mask_tensor::<B>(history, device);
982        let history_hidden_cond = self.embed_history(tokens.clone(), tokens_mask.clone(), false);
983        let mut history_hidden = if use_cfg {
984            let history_hidden_uncond = self.embed_history(tokens, tokens_mask, true);
985            Tensor::cat(vec![history_hidden_cond, history_hidden_uncond], 0)
986        } else {
987            history_hidden_cond
988        };
989
990        if let Some(index) = muq_insert_index {
991            let muq_zero = Tensor::<B, 2>::zeros([1, HEARTMULA_MUQ_DIM], device);
992            let muq_hidden =
993                self.muq_linear
994                    .forward(muq_zero)
995                    .reshape([1, 1, HEARTMULA_HIDDEN_SIZE]);
996            history_hidden = if use_cfg {
997                let uncond_hidden = self
998                    .unconditional_text_embedding
999                    .forward(Tensor::<B, 2, Int>::zeros([1, 1], device))
1000                    .reshape([1, 1, HEARTMULA_HIDDEN_SIZE]);
1001                let replacement = Tensor::cat(vec![muq_hidden, uncond_hidden], 0);
1002                splice_sequence_token(history_hidden, replacement, index)
1003            } else {
1004                splice_sequence_token(history_hidden, muq_hidden, index)
1005            };
1006        }
1007
1008        let positions = position_tensor::<B>((0..history.len() as i64).collect(), device);
1009        self.backbone
1010            .forward_prefill(history_hidden, positions, cache)
1011    }
1012
1013    fn decode_frame_from_last_hidden(
1014        &self,
1015        device: &B::Device,
1016        last_hidden: Tensor<B, 2>,
1017        temperature: f32,
1018        topk: usize,
1019        cfg_scale: f32,
1020    ) -> Result<Vec<i64>> {
1021        let use_cfg = cfg_scale > 1.0;
1022
1023        let codebook0_logits = self.codebook0_head.forward(last_hidden.clone());
1024
1025        let cond_codebook0_logits = if use_cfg {
1026            codebook0_logits
1027                .clone()
1028                .slice([0..1, 0..self.audio_vocab_size()])
1029        } else {
1030            codebook0_logits.clone()
1031        };
1032        let uncond_codebook0_logits = if use_cfg {
1033            codebook0_logits
1034                .clone()
1035                .slice([1..2, 0..self.audio_vocab_size()])
1036        } else {
1037            codebook0_logits.clone()
1038        };
1039        let codebook0_logits = if use_cfg {
1040            uncond_codebook0_logits.clone()
1041                + (cond_codebook0_logits - uncond_codebook0_logits) * cfg_scale
1042        } else {
1043            codebook0_logits
1044        };
1045
1046        let mut frame = Vec::with_capacity(HEARTMULA_AUDIO_CODEBOOKS);
1047        let first_token = sample_token(&codebook0_logits, temperature, topk)?;
1048        frame.push(first_token);
1049
1050        let mut decoder_cache = self.decoder.new_cache();
1051        let c0_embed = self.embed_audio_token(device, 0, first_token);
1052        let c0_embed = if use_cfg {
1053            Tensor::cat(vec![c0_embed.clone(), c0_embed], 0)
1054        } else {
1055            c0_embed
1056        };
1057        let decoder_input = Tensor::cat(
1058            vec![last_hidden.clone().unsqueeze_dim(1), c0_embed.clone()],
1059            1,
1060        );
1061        let decoder_input = self.projection.forward(decoder_input);
1062        let first_decoder_h = self.decoder.forward_prefill(
1063            decoder_input,
1064            position_tensor::<B>(vec![0, 1], device),
1065            &mut decoder_cache,
1066        )?;
1067        let mut current_embed: Option<Tensor<B, 3>> = None;
1068        let mut next_decoder_pos = 2_i64;
1069        for codebook in 1..HEARTMULA_AUDIO_CODEBOOKS {
1070            let head = self
1071                .audio_head
1072                .val()
1073                .slice([
1074                    codebook - 1..codebook,
1075                    0..HEARTMULA_HIDDEN_SIZE,
1076                    0..self.audio_head.dims()[2],
1077                ])
1078                .reshape([HEARTMULA_HIDDEN_SIZE, self.audio_head.dims()[2]]);
1079            let logits = if codebook == 1 {
1080                first_decoder_h.clone().matmul(head.clone())
1081            } else {
1082                let embed = current_embed
1083                    .clone()
1084                    .ok_or_else(|| anyhow!("missing decoder embed for codebook {}", codebook))?;
1085                let decoder_input = self.projection.forward(embed);
1086                let last_decoder_h = self.decoder.forward_incremental(
1087                    decoder_input,
1088                    single_position_tensor::<B>(next_decoder_pos, device),
1089                    &mut decoder_cache,
1090                )?;
1091                next_decoder_pos += 1;
1092                last_decoder_h.matmul(head.clone())
1093            };
1094            let logits = if use_cfg {
1095                let cond_logits = logits.clone().slice([0..1, 0..self.audio_vocab_size()]);
1096                let uncond_logits = logits.slice([1..2, 0..self.audio_vocab_size()]);
1097                uncond_logits.clone() + (cond_logits - uncond_logits) * cfg_scale
1098            } else {
1099                logits
1100            };
1101
1102            let token = sample_token(&logits, temperature, topk)?;
1103            frame.push(token);
1104            current_embed = Some(self.embed_audio_token(device, codebook, token));
1105            if use_cfg {
1106                let embed = current_embed
1107                    .clone()
1108                    .ok_or_else(|| anyhow!("missing decoder embed after sampling"))?;
1109                current_embed = Some(Tensor::cat(vec![embed.clone(), embed], 0));
1110            }
1111        }
1112
1113        Ok(frame)
1114    }
1115
1116    fn embed_history(
1117        &self,
1118        tokens: Tensor<B, 3, Int>,
1119        tokens_mask: Tensor<B, 3, Bool>,
1120        use_unconditional_text: bool,
1121    ) -> Tensor<B, 3> {
1122        let [batch, seq_len, _] = tokens.dims();
1123        let text_ids = tokens
1124            .clone()
1125            .slice([
1126                0..batch,
1127                0..seq_len,
1128                HEARTMULA_AUDIO_CODEBOOKS..HEARTMULA_PARALLEL_TOKENS,
1129            ])
1130            .reshape([batch, seq_len]);
1131        let audio_ids = tokens
1132            .slice([0..batch, 0..seq_len, 0..HEARTMULA_AUDIO_CODEBOOKS])
1133            .reshape([batch, seq_len * HEARTMULA_AUDIO_CODEBOOKS]);
1134        let offsets = (0..HEARTMULA_AUDIO_CODEBOOKS)
1135            .map(|index| (index * self.audio_vocab_size()) as i64)
1136            .collect::<Vec<_>>();
1137        let offset_tensor =
1138            Tensor::<B, 1, Int>::from_data(offsets.as_slice(), &tokens_mask.device()).reshape([
1139                1,
1140                1,
1141                HEARTMULA_AUDIO_CODEBOOKS,
1142            ]);
1143        let shifted_audio_ids = audio_ids
1144            .reshape([batch, seq_len, HEARTMULA_AUDIO_CODEBOOKS])
1145            .add(offset_tensor);
1146
1147        let text_embeds = if use_unconditional_text {
1148            self.unconditional_text_embedding
1149                .forward(Tensor::<B, 2, Int>::zeros(
1150                    [batch, seq_len],
1151                    &tokens_mask.device(),
1152                ))
1153                .unsqueeze_dim(2)
1154        } else {
1155            self.text_embeddings.forward(text_ids).unsqueeze_dim(2)
1156        };
1157        let audio_embeds = self.audio_embeddings.forward(shifted_audio_ids);
1158        let text_embeds = if text_embeds.dims()[0] == audio_embeds.dims()[0] {
1159            text_embeds
1160        } else {
1161            text_embeds.repeat_dim(0, audio_embeds.dims()[0])
1162        };
1163        let embeds = Tensor::cat(vec![audio_embeds, text_embeds], 2);
1164        let mask = tokens_mask
1165            .reshape([batch, seq_len, HEARTMULA_PARALLEL_TOKENS, 1])
1166            .repeat_dim(3, HEARTMULA_HIDDEN_SIZE)
1167            .float();
1168
1169        (embeds * mask)
1170            .sum_dim(2)
1171            .reshape([batch, seq_len, HEARTMULA_HIDDEN_SIZE])
1172    }
1173
1174    fn embed_audio_token(&self, device: &B::Device, codebook: usize, token: i64) -> Tensor<B, 3> {
1175        let offset_token = token + (codebook * self.audio_vocab_size()) as i64;
1176        self.audio_embeddings
1177            .embed_offset_token(device, offset_token)
1178    }
1179
1180    fn audio_vocab_size(&self) -> usize {
1181        self.audio_head.dims()[2]
1182    }
1183
1184    fn embed_single_history_row(
1185        &self,
1186        device: &B::Device,
1187        row: &[i64; HEARTMULA_PARALLEL_TOKENS],
1188    ) -> Tensor<B, 3> {
1189        let tokens = Tensor::<B, 3, Int>::from_data(
1190            TensorData::new(row.to_vec(), [1, 1, HEARTMULA_PARALLEL_TOKENS]),
1191            device,
1192        );
1193        let mask = Tensor::<B, 3, Bool>::from_data(
1194            TensorData::new(
1195                vec![true, true, true, true, true, true, true, true, false],
1196                [1, 1, HEARTMULA_PARALLEL_TOKENS],
1197            ),
1198            device,
1199        );
1200        self.embed_history(tokens, mask, false)
1201    }
1202}
1203
1204fn sync_and_cleanup_backend<B: Backend>(device: &B::Device) -> Result<()> {
1205    B::sync(device)?;
1206    B::memory_cleanup(device);
1207    Ok(())
1208}
1209
1210impl<B: Backend> SplitAudioEmbeddings<B> {
1211    fn new_placeholder(vocab_size: usize) -> Self {
1212        Self {
1213            table: None,
1214            vocab_size,
1215        }
1216    }
1217
1218    fn load_from_data(device: &B::Device, data: TensorData, vocab_size: usize) -> Result<Self> {
1219        let shape = data.shape.clone();
1220        let expected_rows = vocab_size * HEARTMULA_AUDIO_CODEBOOKS;
1221        if shape.as_slice() != [expected_rows, HEARTMULA_HIDDEN_SIZE] {
1222            anyhow::bail!(
1223                "unexpected audio_embeddings.weight shape {:?}, expected [{}, {}]",
1224                shape,
1225                expected_rows,
1226                HEARTMULA_HIDDEN_SIZE
1227            );
1228        }
1229        Ok(Self {
1230            table: Some(Tensor::<B, 2>::from_data(data, device)),
1231            vocab_size,
1232        })
1233    }
1234
1235    fn forward(&self, offset_audio_ids: Tensor<B, 3, Int>) -> Tensor<B, 4> {
1236        let [batch, seq_len, codebooks] = offset_audio_ids.dims();
1237        debug_assert_eq!(codebooks, HEARTMULA_AUDIO_CODEBOOKS);
1238        let table = self
1239            .table
1240            .as_ref()
1241            .expect("audio embeddings must be loaded before use")
1242            .clone();
1243        let ids = offset_audio_ids.reshape([batch * seq_len * codebooks]);
1244        table
1245            .select(0, ids)
1246            .reshape([batch, seq_len, codebooks, HEARTMULA_HIDDEN_SIZE])
1247    }
1248
1249    fn embed_offset_token(&self, device: &B::Device, offset_token: i64) -> Tensor<B, 3> {
1250        debug_assert!((offset_token as usize) < self.vocab_size * HEARTMULA_AUDIO_CODEBOOKS);
1251        let ids = Tensor::<B, 1, Int>::from_data([offset_token], device);
1252        self.table
1253            .as_ref()
1254            .expect("audio embeddings must be loaded before use")
1255            .clone()
1256            .select(0, ids)
1257            .reshape([1, 1, HEARTMULA_HIDDEN_SIZE])
1258    }
1259}
1260
1261impl<B: Backend> Module<B> for SplitAudioEmbeddings<B> {
1262    type Record = EmptyRecord;
1263
1264    fn visit<V: ModuleVisitor<B>>(&self, _visitor: &mut V) {}
1265
1266    fn map<M: ModuleMapper<B>>(self, _mapper: &mut M) -> Self {
1267        self
1268    }
1269
1270    fn load_record(self, _record: Self::Record) -> Self {
1271        self
1272    }
1273
1274    fn into_record(self) -> Self::Record {
1275        EmptyRecord::new()
1276    }
1277
1278    fn to_device(self, device: &B::Device) -> Self {
1279        Self {
1280            table: self.table.map(|tensor| tensor.to_device(device)),
1281            vocab_size: self.vocab_size,
1282        }
1283    }
1284
1285    fn fork(self, device: &B::Device) -> Self {
1286        Self {
1287            table: self.table.map(|tensor| tensor.fork(device)),
1288            vocab_size: self.vocab_size,
1289        }
1290    }
1291
1292    fn collect_devices(&self, mut devices: Devices<B>) -> Devices<B> {
1293        if let Some(tensor) = &self.table {
1294            let device = tensor.device();
1295            if !devices.contains(&device) {
1296                devices.push(device);
1297            }
1298        }
1299        devices
1300    }
1301}
1302
1303impl<B: Backend> ModuleDisplayDefault for SplitAudioEmbeddings<B> {
1304    fn content(&self, content: Content) -> Option<Content> {
1305        content
1306            .add("table_loaded", &self.table.is_some())
1307            .add("vocab_size", &self.vocab_size)
1308            .optional()
1309    }
1310}
1311
1312impl<B: Backend> ModuleDisplay for SplitAudioEmbeddings<B> {}
1313
1314impl<B: AutodiffBackend> AutodiffModule<B> for SplitAudioEmbeddings<B> {
1315    type InnerModule = SplitAudioEmbeddings<B::InnerBackend>;
1316
1317    fn valid(&self) -> Self::InnerModule {
1318        SplitAudioEmbeddings {
1319            table: self.table.as_ref().map(|tensor| tensor.valid()),
1320            vocab_size: self.vocab_size,
1321        }
1322    }
1323
1324    fn from_inner(module: Self::InnerModule) -> Self {
1325        SplitAudioEmbeddings {
1326            table: module
1327                .table
1328                .map(|tensor| Tensor::<B, 2>::from_data(tensor.to_data(), &tensor.device())),
1329            vocab_size: module.vocab_size,
1330        }
1331    }
1332}
1333
1334impl<B: Backend> HeartmulaTransformer<B> {
1335    fn new(device: &B::Device, layer_count: usize, num_heads: usize, num_kv_heads: usize) -> Self {
1336        let layers = (0..layer_count)
1337            .map(|_| HeartmulaTransformerLayer::new(device, num_heads, num_kv_heads))
1338            .collect();
1339        Self {
1340            layers,
1341            norm: HeartmulaRmsNorm::new(device, HEARTMULA_HIDDEN_SIZE, HEARTMULA_NORM_EPSILON),
1342        }
1343    }
1344
1345    fn new_cache(&self) -> HeartmulaTransformerCache<B> {
1346        HeartmulaTransformerCache {
1347            layers: (0..self.layers.len())
1348                .map(|_| HeartmulaAttentionCache {
1349                    key: None,
1350                    value: None,
1351                })
1352                .collect(),
1353        }
1354    }
1355
1356    fn forward_incremental(
1357        &self,
1358        mut hidden: Tensor<B, 3>,
1359        position: Tensor<B, 2, Int>,
1360        cache: &mut HeartmulaTransformerCache<B>,
1361    ) -> Result<Tensor<B, 2>> {
1362        for (layer, layer_cache) in self.layers.iter().zip(cache.layers.iter_mut()) {
1363            hidden = layer.forward_incremental(hidden, position.clone(), layer_cache)?;
1364        }
1365        Ok(take_last_token(self.norm.forward(hidden)))
1366    }
1367
1368    fn forward_prefill(
1369        &self,
1370        mut hidden: Tensor<B, 3>,
1371        positions: Tensor<B, 2, Int>,
1372        cache: &mut HeartmulaTransformerCache<B>,
1373    ) -> Result<Tensor<B, 2>> {
1374        for (layer, layer_cache) in self.layers.iter().zip(cache.layers.iter_mut()) {
1375            hidden = layer.forward_prefill(hidden, positions.clone(), layer_cache)?;
1376        }
1377        Ok(take_last_token(self.norm.forward(hidden)))
1378    }
1379}
1380
1381impl<B: Backend> HeartmulaTransformerLayer<B> {
1382    fn new(device: &B::Device, num_heads: usize, num_kv_heads: usize) -> Self {
1383        Self {
1384            attn: HeartmulaAttention::new(device, num_heads, num_kv_heads),
1385            mlp: HeartmulaMlp::new(device),
1386            sa_norm: HeartmulaRmsNorm::new(device, HEARTMULA_HIDDEN_SIZE, HEARTMULA_NORM_EPSILON),
1387            mlp_norm: HeartmulaRmsNorm::new(device, HEARTMULA_HIDDEN_SIZE, HEARTMULA_NORM_EPSILON),
1388        }
1389    }
1390
1391    fn forward_incremental(
1392        &self,
1393        hidden: Tensor<B, 3>,
1394        position: Tensor<B, 2, Int>,
1395        cache: &mut HeartmulaAttentionCache<B>,
1396    ) -> Result<Tensor<B, 3>> {
1397        let attn_hidden =
1398            self.attn
1399                .forward_incremental(self.sa_norm.forward(hidden.clone()), position, cache)?;
1400        let hidden = hidden + attn_hidden;
1401        let mlp_hidden = self.mlp.forward(self.mlp_norm.forward(hidden.clone()));
1402        Ok(hidden + mlp_hidden)
1403    }
1404
1405    fn forward_prefill(
1406        &self,
1407        hidden: Tensor<B, 3>,
1408        positions: Tensor<B, 2, Int>,
1409        cache: &mut HeartmulaAttentionCache<B>,
1410    ) -> Result<Tensor<B, 3>> {
1411        let attn_hidden =
1412            self.attn
1413                .forward_prefill(self.sa_norm.forward(hidden.clone()), positions, cache)?;
1414        let hidden = hidden + attn_hidden;
1415        let mlp_hidden = self.mlp.forward(self.mlp_norm.forward(hidden.clone()));
1416        Ok(hidden + mlp_hidden)
1417    }
1418}
1419
1420impl<B: Backend> HeartmulaAttention<B> {
1421    fn new(device: &B::Device, num_heads: usize, num_kv_heads: usize) -> Self {
1422        let head_dim = HEARTMULA_HIDDEN_SIZE / num_heads;
1423        Self {
1424            q_proj: linear_no_bias(device, HEARTMULA_HIDDEN_SIZE, num_heads * head_dim),
1425            k_proj: linear_no_bias(device, HEARTMULA_HIDDEN_SIZE, num_kv_heads * head_dim),
1426            v_proj: linear_no_bias(device, HEARTMULA_HIDDEN_SIZE, num_kv_heads * head_dim),
1427            output_proj: linear_no_bias(device, HEARTMULA_HIDDEN_SIZE, HEARTMULA_HIDDEN_SIZE),
1428            meta: AttentionMeta {
1429                num_heads,
1430                num_kv_heads,
1431                head_dim,
1432            },
1433        }
1434    }
1435
1436    fn forward_incremental(
1437        &self,
1438        hidden: Tensor<B, 3>,
1439        position: Tensor<B, 2, Int>,
1440        cache: &mut HeartmulaAttentionCache<B>,
1441    ) -> Result<Tensor<B, 3>> {
1442        let [batch, seq_len, _] = hidden.dims();
1443        let q = self.q_proj.forward(hidden.clone()).reshape([
1444            batch,
1445            seq_len,
1446            self.meta.num_heads,
1447            self.meta.head_dim,
1448        ]);
1449        let k = self.k_proj.forward(hidden.clone()).reshape([
1450            batch,
1451            seq_len,
1452            self.meta.num_kv_heads,
1453            self.meta.head_dim,
1454        ]);
1455        let v = self.v_proj.forward(hidden).reshape([
1456            batch,
1457            seq_len,
1458            self.meta.num_kv_heads,
1459            self.meta.head_dim,
1460        ]);
1461
1462        let q = apply_scaled_rope(q, &position).swap_dims(1, 2);
1463        let k = apply_scaled_rope(k, &position).swap_dims(1, 2);
1464        let v = v.swap_dims(1, 2);
1465
1466        let full_k = if let Some(previous) = &cache.key {
1467            Tensor::cat(vec![previous.clone(), k], 2)
1468        } else {
1469            k
1470        };
1471        let full_v = if let Some(previous) = &cache.value {
1472            Tensor::cat(vec![previous.clone(), v], 2)
1473        } else {
1474            v
1475        };
1476        cache.key = Some(full_k.clone());
1477        cache.value = Some(full_v.clone());
1478
1479        let (full_k_for_attn, full_v_for_attn) = if self.meta.num_heads != self.meta.num_kv_heads {
1480            let repeats = self.meta.num_heads / self.meta.num_kv_heads;
1481            (
1482                repeat_cached_kv_heads(full_k, repeats),
1483                repeat_cached_kv_heads(full_v, repeats),
1484            )
1485        } else {
1486            (full_k, full_v)
1487        };
1488
1489        let weights = softmax(
1490            q.matmul(full_k_for_attn.swap_dims(2, 3))
1491                .mul_scalar(1.0 / (self.meta.head_dim as f32).sqrt()),
1492            3,
1493        );
1494        let attended = weights.matmul(full_v_for_attn).swap_dims(1, 2).reshape([
1495            batch,
1496            seq_len,
1497            HEARTMULA_HIDDEN_SIZE,
1498        ]);
1499        Ok(self.output_proj.forward(attended))
1500    }
1501
1502    fn forward_prefill(
1503        &self,
1504        hidden: Tensor<B, 3>,
1505        positions: Tensor<B, 2, Int>,
1506        cache: &mut HeartmulaAttentionCache<B>,
1507    ) -> Result<Tensor<B, 3>> {
1508        let [batch, seq_len, _] = hidden.dims();
1509        let q = self.q_proj.forward(hidden.clone()).reshape([
1510            batch,
1511            seq_len,
1512            self.meta.num_heads,
1513            self.meta.head_dim,
1514        ]);
1515        let k = self.k_proj.forward(hidden.clone()).reshape([
1516            batch,
1517            seq_len,
1518            self.meta.num_kv_heads,
1519            self.meta.head_dim,
1520        ]);
1521        let v = self.v_proj.forward(hidden).reshape([
1522            batch,
1523            seq_len,
1524            self.meta.num_kv_heads,
1525            self.meta.head_dim,
1526        ]);
1527
1528        let q = apply_scaled_rope(q, &positions).swap_dims(1, 2);
1529        let k = apply_scaled_rope(k, &positions).swap_dims(1, 2);
1530        let v = v.swap_dims(1, 2);
1531        cache.key = Some(k.clone());
1532        cache.value = Some(v.clone());
1533
1534        let (k_for_attn, v_for_attn) = if self.meta.num_heads != self.meta.num_kv_heads {
1535            let repeats = self.meta.num_heads / self.meta.num_kv_heads;
1536            (
1537                repeat_cached_kv_heads(k, repeats),
1538                repeat_cached_kv_heads(v, repeats),
1539            )
1540        } else {
1541            (k, v)
1542        };
1543
1544        let scores = q
1545            .matmul(k_for_attn.clone().swap_dims(2, 3))
1546            .mul_scalar(1.0 / (self.meta.head_dim as f32).sqrt());
1547        let mask = causal_mask::<B>(seq_len, &scores.device());
1548        let weights = softmax(scores.mask_fill(mask, -1.0e9), 3);
1549        let attended = weights.matmul(v_for_attn).swap_dims(1, 2).reshape([
1550            batch,
1551            seq_len,
1552            HEARTMULA_HIDDEN_SIZE,
1553        ]);
1554        Ok(self.output_proj.forward(attended))
1555    }
1556}
1557
1558impl<B: Backend> HeartmulaMlp<B> {
1559    fn new(device: &B::Device) -> Self {
1560        Self {
1561            w1: linear_no_bias(device, HEARTMULA_HIDDEN_SIZE, HEARTMULA_MLP_DIM),
1562            w2: linear_no_bias(device, HEARTMULA_MLP_DIM, HEARTMULA_HIDDEN_SIZE),
1563            w3: linear_no_bias(device, HEARTMULA_HIDDEN_SIZE, HEARTMULA_MLP_DIM),
1564        }
1565    }
1566
1567    fn forward(&self, hidden: Tensor<B, 3>) -> Tensor<B, 3> {
1568        let gate = silu(self.w1.forward(hidden.clone()));
1569        let up = self.w3.forward(hidden);
1570        self.w2.forward(gate * up)
1571    }
1572}
1573
1574impl<B: Backend> HeartmulaRmsNorm<B> {
1575    fn new(device: &B::Device, hidden_size: usize, epsilon: f64) -> Self {
1576        Self {
1577            scale: Param::from_tensor(Tensor::<B, 1>::ones([hidden_size], device)),
1578            epsilon,
1579        }
1580    }
1581
1582    fn forward<const D: usize>(&self, hidden: Tensor<B, D>) -> Tensor<B, D> {
1583        let dtype = hidden.dtype();
1584        let rms = (hidden.clone().cast(DType::F32).square().mean_dim(D - 1) + self.epsilon).sqrt();
1585        (hidden / rms.cast(dtype)) * self.scale.val().unsqueeze()
1586    }
1587}
1588
1589pub fn tokenize_text(tokenizer_json: &Path, text: &str) -> Result<Vec<i64>> {
1590    let tokenizer = Tokenizer::from_json(tokenizer_json).map_err(|e| {
1591        anyhow!(
1592            "failed to load tokenizer from {}: {e}",
1593            tokenizer_json.display()
1594        )
1595    })?;
1596    let encoding = tokenizer.encode(text, true);
1597    Ok(encoding.ids.into_iter().map(i64::from).collect())
1598}
1599
1600pub fn default_tags() -> &'static str {
1601    "<tag></tag>"
1602}
1603
1604pub fn normalize_tags(tags: &str) -> String {
1605    let mut normalized = tags.trim().to_lowercase();
1606
1607    while normalized.contains(", ") {
1608        normalized = normalized.replace(", ", ",");
1609    }
1610    if !normalized.starts_with("<tag>") {
1611        normalized = format!("<tag>{normalized}");
1612    }
1613    if !normalized.ends_with("</tag>") {
1614        normalized.push_str("</tag>");
1615    }
1616    normalized
1617}
1618
1619pub fn write_frames_json(path: &Path, lyrics: &str, tags: &str, frames: &[Vec<i64>]) -> Result<()> {
1620    let payload = HeartmulaJsonOutput {
1621        model: "heartmula".to_string(),
1622        runtime: "burn-token-generator".to_string(),
1623        tags: tags.to_owned(),
1624        lyrics: lyrics.to_owned(),
1625        frames: frames.to_vec(),
1626        frame_count: frames.len(),
1627        sample_rate_hz: 48_000,
1628    };
1629    std::fs::write(path, serde_json::to_vec_pretty(&payload)?)
1630        .with_context(|| format!("failed to write {}", path.display()))
1631}
1632
1633#[allow(clippy::too_many_arguments)]
1634pub fn decode_frames_to_wav<B: burn::prelude::Backend>(
1635    model_dir: &Path,
1636    _backend_arg: &str,
1637    float_size_arg: &str,
1638    frames_json: &Path,
1639    output_wav: &Path,
1640    duration_seconds: f32,
1641    device: &B::Device,
1642    ode_steps: usize,
1643    decoder_seed: u64,
1644) -> Result<()> {
1645    if let Some(stage) =
1646        env::var_os(HEARTCODEC_STAGE_ENV).and_then(|value| value.into_string().ok())
1647    {
1648        return match stage.as_str() {
1649            HEARTCODEC_STAGE_FLOW => decode_frames_to_plan_rust::<B>(
1650                model_dir,
1651                frames_json,
1652                device,
1653                ode_steps,
1654                &prepare_shared_decoder_initial_latent(frames_json, decoder_seed, output_wav)?,
1655            ),
1656            HEARTCODEC_STAGE_SCALAR => decode_plan_to_wav_rust::<B>(model_dir, output_wav, device),
1657            other => Err(anyhow!("unsupported HeartCodec stage '{other}'")),
1658        };
1659    }
1660
1661    let _ = float_size_arg;
1662    let _ = duration_seconds;
1663    decode_frames_to_wav_rust::<B>(
1664        model_dir,
1665        frames_json,
1666        output_wav,
1667        duration_seconds,
1668        decoder_seed,
1669        device,
1670        ode_steps,
1671    )
1672}
1673
1674fn resolve_heartcodec_burnpack_path(model_dir: &Path) -> PathBuf {
1675    model_dir.join("heartcodec.bpk")
1676}
1677
1678fn decode_frames_to_wav_rust<B: burn::prelude::Backend>(
1679    model_dir: &Path,
1680    frames_json: &Path,
1681    output_wav: &Path,
1682    duration_seconds: f32,
1683    decoder_seed: u64,
1684    device: &B::Device,
1685    ode_steps: usize,
1686) -> Result<()> {
1687    let frames_text = std::fs::read_to_string(frames_json)
1688        .with_context(|| format!("failed to read {}", frames_json.display()))?;
1689    let _payload: HeartmulaJsonOutput = serde_json::from_str(&frames_text)
1690        .with_context(|| format!("failed to parse {}", frames_json.display()))?;
1691    B::seed(device, 0);
1692    let initial_latent_json =
1693        prepare_shared_decoder_initial_latent(frames_json, decoder_seed, output_wav)?;
1694    let stage_plan_json = output_wav.with_extension("heartcodec-stage-plan.bin");
1695
1696    unsafe {
1697        std::env::set_var(HEARTCODEC_STAGE_PLAN_JSON_ENV, &stage_plan_json);
1698    }
1699
1700    let plan_result = decode_frames_to_plan_rust::<B>(
1701        model_dir,
1702        frames_json,
1703        device,
1704        ode_steps,
1705        &initial_latent_json,
1706    );
1707    sync_and_cleanup_backend::<B>(device)?;
1708    plan_result?;
1709
1710    let decode_result = decode_plan_to_wav_rust::<B>(model_dir, output_wav, device);
1711    sync_and_cleanup_backend::<B>(device)?;
1712
1713    let _ = std::fs::remove_file(&stage_plan_json);
1714    let _ = std::fs::remove_file(&initial_latent_json);
1715
1716    decode_result?;
1717    let _ = duration_seconds;
1718    Ok(())
1719}
1720
1721pub fn decode_frames_to_plan_rust<B: burn::prelude::Backend>(
1722    model_dir: &Path,
1723    frames_json: &Path,
1724    device: &B::Device,
1725    ode_steps: usize,
1726    initial_latent_json: &Path,
1727) -> Result<()> {
1728    let frames_text = std::fs::read_to_string(frames_json)
1729        .with_context(|| format!("failed to read {}", frames_json.display()))?;
1730    let payload: HeartmulaJsonOutput = serde_json::from_str(&frames_text)
1731        .with_context(|| format!("failed to parse {}", frames_json.display()))?;
1732    let frames = payload.frames;
1733    B::seed(device, 0);
1734    let codes = frames_to_tensor::<B>(&frames, device);
1735    let codec_path = resolve_heartcodec_burnpack_path(model_dir);
1736    let flow_matching =
1737        crate::heartcodec::FlowMatching::<B>::load_from_burnpack(&codec_path, device)?;
1738    let initial_latent = load_initial_latent_tensor::<B>(initial_latent_json, device)?;
1739    let plan = crate::heartcodec::HeartCodecModel::<B>::build_scalar_decode_plan_impl(
1740        &flow_matching,
1741        1.25,
1742        ode_steps,
1743        codes,
1744        initial_latent,
1745    );
1746    let stage_plan_json = current_codec_stage_plan_json()?;
1747    save_codec_stage_plan(&stage_plan_json, plan)?;
1748    Ok(())
1749}
1750
1751pub fn decode_plan_to_wav_rust<B: burn::prelude::Backend>(
1752    model_dir: &Path,
1753    output_wav: &Path,
1754    device: &B::Device,
1755) -> Result<()> {
1756    let stage_plan_json = current_codec_stage_plan_json()?;
1757    let plan = load_codec_stage_plan::<B>(&stage_plan_json, device)?;
1758    let codec_path = resolve_heartcodec_burnpack_path(model_dir);
1759    let scalar_model = crate::heartcodec::ScalarModel::<B>::from_burnpack(&codec_path, device)?;
1760    let wav = crate::heartcodec::HeartCodecModel::<B>::decode_scalar_plan_impl(&scalar_model, plan);
1761    write_decoder_wav(output_wav, wav, 0.0)
1762}
1763
1764fn write_decoder_wav<B: burn::prelude::Backend>(
1765    output_wav: &Path,
1766    wav: Tensor<B, 3>,
1767    _duration_seconds: f32,
1768) -> Result<()> {
1769    let dims = wav.dims();
1770    let samples: Vec<f32> = wav.cast(DType::F32).to_data().to_vec::<f32>()?;
1771    match dims.as_slice() {
1772        [channels, 1, frames] if *channels > 1 => {
1773            crate::heartcodec::write_wav_from_f32_interleaved(
1774                &samples, *channels, *frames, 48_000, output_wav,
1775            )
1776        }
1777        [1, channels, frames] if *channels > 1 => {
1778            crate::heartcodec::write_wav_from_f32_interleaved(
1779                &samples, *channels, *frames, 48_000, output_wav,
1780            )
1781        }
1782        [1, 1, frames] => {
1783            crate::heartcodec::write_wav_from_f32(&samples[..*frames], 48_000, output_wav)
1784        }
1785        _ => crate::heartcodec::write_wav_from_f32(&samples, 48_000, output_wav),
1786    }
1787}
1788
1789fn current_codec_stage_plan_json() -> Result<PathBuf> {
1790    env::var_os(HEARTCODEC_STAGE_PLAN_JSON_ENV)
1791        .map(PathBuf::from)
1792        .ok_or_else(|| {
1793            anyhow!("missing {HEARTCODEC_STAGE_PLAN_JSON_ENV} for staged HeartCodec decode")
1794        })
1795}
1796
1797fn save_codec_stage_plan<B: burn::prelude::Backend>(
1798    path: &Path,
1799    plan: crate::heartcodec::ScalarDecodePlan<B>,
1800) -> Result<()> {
1801    let file =
1802        File::create(path).with_context(|| format!("failed to create {}", path.display()))?;
1803    let mut writer = BufWriter::new(file);
1804    writer
1805        .write_all(HEARTCODEC_STAGE_PLAN_MAGIC)
1806        .with_context(|| format!("failed to write {}", path.display()))?;
1807    write_u64(&mut writer, plan.target_len)?;
1808    write_u64(&mut writer, plan.audio_target_len)?;
1809    write_u64(&mut writer, plan.windows.len())?;
1810    for window in plan.windows {
1811        let dims = window.dims();
1812        let data = window.cast(DType::F32).to_data().to_vec::<f32>()?;
1813        write_dims(&mut writer, dims)?;
1814        write_f32_slice(&mut writer, &data)?;
1815    }
1816    writer
1817        .flush()
1818        .with_context(|| format!("failed to flush {}", path.display()))
1819}
1820
1821fn load_codec_stage_plan<B: burn::prelude::Backend>(
1822    path: &Path,
1823    device: &B::Device,
1824) -> Result<crate::heartcodec::ScalarDecodePlan<B>> {
1825    let file = File::open(path).with_context(|| format!("failed to open {}", path.display()))?;
1826    let mut reader = BufReader::new(file);
1827    let mut magic = [0_u8; HEARTCODEC_STAGE_PLAN_MAGIC.len()];
1828    reader
1829        .read_exact(&mut magic)
1830        .with_context(|| format!("failed to read {}", path.display()))?;
1831    if &magic != HEARTCODEC_STAGE_PLAN_MAGIC {
1832        anyhow::bail!("invalid HeartCodec stage plan format in {}", path.display());
1833    }
1834    let target_len = read_u64(&mut reader)? as usize;
1835    let audio_target_len = read_u64(&mut reader)? as usize;
1836    let window_count = read_u64(&mut reader)? as usize;
1837    let mut windows = Vec::with_capacity(window_count);
1838    for _ in 0..window_count {
1839        let dims = read_dims(&mut reader)?;
1840        let data = read_f32_vec(&mut reader)?;
1841        windows.push(Tensor::<B, 3>::from_data(
1842            TensorData::new(data, dims),
1843            device,
1844        ));
1845    }
1846    Ok(crate::heartcodec::ScalarDecodePlan {
1847        target_len,
1848        audio_target_len,
1849        windows,
1850    })
1851}
1852
1853fn write_u64(writer: &mut dyn Write, value: usize) -> Result<()> {
1854    writer.write_all(&(value as u64).to_le_bytes())?;
1855    Ok(())
1856}
1857
1858fn read_u64(reader: &mut dyn Read) -> Result<u64> {
1859    let mut bytes = [0_u8; 8];
1860    reader.read_exact(&mut bytes)?;
1861    Ok(u64::from_le_bytes(bytes))
1862}
1863
1864fn write_dims(writer: &mut dyn Write, dims: [usize; 3]) -> Result<()> {
1865    for value in dims {
1866        write_u64(writer, value)?;
1867    }
1868    Ok(())
1869}
1870
1871fn read_dims(reader: &mut dyn Read) -> Result<[usize; 3]> {
1872    Ok([
1873        read_u64(reader)? as usize,
1874        read_u64(reader)? as usize,
1875        read_u64(reader)? as usize,
1876    ])
1877}
1878
1879fn write_f32_slice(writer: &mut dyn Write, values: &[f32]) -> Result<()> {
1880    write_u64(writer, values.len())?;
1881    let mut bytes = vec![0_u8; std::mem::size_of_val(values)];
1882    bytes
1883        .par_chunks_mut(std::mem::size_of::<f32>())
1884        .zip(values.par_iter())
1885        .for_each(|(chunk, value)| chunk.copy_from_slice(&value.to_le_bytes()));
1886    writer.write_all(&bytes)?;
1887    Ok(())
1888}
1889
1890fn read_f32_vec(reader: &mut dyn Read) -> Result<Vec<f32>> {
1891    let len = read_u64(reader)? as usize;
1892    let mut bytes = vec![0_u8; len * std::mem::size_of::<f32>()];
1893    reader.read_exact(&mut bytes)?;
1894    let mut values = vec![0.0_f32; len];
1895    values
1896        .par_iter_mut()
1897        .enumerate()
1898        .for_each(|(index, value)| {
1899            let offset = index * std::mem::size_of::<f32>();
1900            *value = f32::from_le_bytes([
1901                bytes[offset],
1902                bytes[offset + 1],
1903                bytes[offset + 2],
1904                bytes[offset + 3],
1905            ]);
1906        });
1907    Ok(values)
1908}
1909
1910fn prepare_shared_decoder_initial_latent(
1911    _frames_json: &Path,
1912    decoder_seed: u64,
1913    output_wav: &Path,
1914) -> Result<PathBuf> {
1915    let latent_length = (HEARTCODEC_SEGMENT_DURATION_SECONDS * 25.0) as usize;
1916    let dims = [1, latent_length, 256];
1917    let data = generate_decoder_latent_data(decoder_seed, dims[0] * dims[1] * dims[2]);
1918    let latent = LatentTensorFile { dims, data };
1919
1920    let stem = output_wav
1921        .file_stem()
1922        .and_then(|value| value.to_str())
1923        .unwrap_or("decoder");
1924    let path = std::env::temp_dir().join(format!(
1925        "maolan-{stem}-decoder-seed-{decoder_seed}-latent-{latent_length}.json"
1926    ));
1927    fs::write(&path, serde_json::to_vec(&latent)?)
1928        .with_context(|| format!("failed to write {}", path.display()))?;
1929    Ok(path)
1930}
1931
1932fn load_initial_latent_tensor<B: burn::prelude::Backend>(
1933    path: &Path,
1934    device: &B::Device,
1935) -> Result<Tensor<B, 3>> {
1936    let text =
1937        fs::read_to_string(path).with_context(|| format!("failed to read {}", path.display()))?;
1938    let payload: LatentTensorFile = serde_json::from_str(&text)
1939        .with_context(|| format!("failed to parse {}", path.display()))?;
1940    Ok(Tensor::<B, 3>::from_data(
1941        TensorData::new(payload.data, payload.dims),
1942        device,
1943    ))
1944}
1945
1946fn generate_decoder_latent_data(seed: u64, len: usize) -> Vec<f32> {
1947    let mut out = Vec::with_capacity(len);
1948    let mut state = seed;
1949    while out.len() < len {
1950        let u1 = uniform01_open(&mut state);
1951        let u2 = uniform01_open(&mut state);
1952        let radius = (-2.0_f64 * u1.ln()).sqrt();
1953        let theta = 2.0_f64 * std::f64::consts::PI * u2;
1954        out.push((radius * theta.cos()) as f32);
1955        if out.len() < len {
1956            out.push((radius * theta.sin()) as f32);
1957        }
1958    }
1959    out
1960}
1961
1962fn uniform01_open(state: &mut u64) -> f64 {
1963    let value = splitmix64_next(state);
1964    let mantissa = (value >> 11) as f64;
1965    ((mantissa + 0.5) / ((1_u64 << 53) as f64)).clamp(f64::MIN_POSITIVE, 1.0 - f64::EPSILON)
1966}
1967
1968fn splitmix64_next(state: &mut u64) -> u64 {
1969    *state = state.wrapping_add(0x9E3779B97F4A7C15);
1970    let mut z = *state;
1971    z = (z ^ (z >> 30)).wrapping_mul(0xBF58476D1CE4E5B9);
1972    z = (z ^ (z >> 27)).wrapping_mul(0x94D049BB133111EB);
1973    z ^ (z >> 31)
1974}
1975
1976fn build_prompt_history(
1977    text_bos_id: i64,
1978    text_eos_id: i64,
1979    lyrics_ids: &[i64],
1980    tags_ids: &[i64],
1981) -> Vec<[i64; HEARTMULA_PARALLEL_TOKENS]> {
1982    let full_tags = normalize_text_ids(text_bos_id, text_eos_id, tags_ids);
1983    let full_lyrics = normalize_text_ids(text_bos_id, text_eos_id, lyrics_ids);
1984
1985    let mut history = Vec::with_capacity(full_tags.len() + 1 + full_lyrics.len());
1986    for token in full_tags {
1987        let mut row = [0_i64; HEARTMULA_PARALLEL_TOKENS];
1988        row[HEARTMULA_AUDIO_CODEBOOKS] = token;
1989        history.push(row);
1990    }
1991    history.push([0_i64; HEARTMULA_PARALLEL_TOKENS]);
1992    for token in full_lyrics {
1993        let mut row = [0_i64; HEARTMULA_PARALLEL_TOKENS];
1994        row[HEARTMULA_AUDIO_CODEBOOKS] = token;
1995        history.push(row);
1996    }
1997    history
1998}
1999
2000fn normalize_text_ids(text_bos_id: i64, text_eos_id: i64, ids: &[i64]) -> Vec<i64> {
2001    let mut normalized = ids.to_vec();
2002    if normalized.first().copied() != Some(text_bos_id) {
2003        normalized.insert(0, text_bos_id);
2004    }
2005    if normalized.last().copied() != Some(text_eos_id) {
2006        normalized.push(text_eos_id);
2007    }
2008    normalized
2009}
2010
2011fn build_audio_history_row(frame: &[i64], empty_id: i64) -> [i64; HEARTMULA_PARALLEL_TOKENS] {
2012    let mut row = [empty_id; HEARTMULA_PARALLEL_TOKENS];
2013    for (index, token) in frame
2014        .iter()
2015        .copied()
2016        .enumerate()
2017        .take(HEARTMULA_AUDIO_CODEBOOKS)
2018    {
2019        row[index] = token;
2020    }
2021    row[HEARTMULA_AUDIO_CODEBOOKS] = empty_id;
2022    row
2023}
2024
2025fn splice_sequence_token<B: Backend>(
2026    hidden: Tensor<B, 3>,
2027    replacement: Tensor<B, 3>,
2028    index: usize,
2029) -> Tensor<B, 3> {
2030    let [batch, seq_len, dim] = hidden.dims();
2031    debug_assert_eq!(replacement.dims(), [batch, 1, dim]);
2032    let mut parts = Vec::new();
2033    if index > 0 {
2034        parts.push(hidden.clone().slice([0..batch, 0..index, 0..dim]));
2035    }
2036    parts.push(replacement);
2037    if index + 1 < seq_len {
2038        parts.push(hidden.slice([0..batch, index + 1..seq_len, 0..dim]));
2039    }
2040    Tensor::cat(parts, 1)
2041}
2042
2043fn history_tokens_tensor<B: Backend>(
2044    history: &[[i64; HEARTMULA_PARALLEL_TOKENS]],
2045    device: &B::Device,
2046) -> Tensor<B, 3, Int> {
2047    let flattened = history
2048        .iter()
2049        .flat_map(|row| row.iter().copied())
2050        .collect::<Vec<_>>();
2051    Tensor::<B, 3, Int>::from_data(
2052        TensorData::new(flattened, [1, history.len(), HEARTMULA_PARALLEL_TOKENS]),
2053        device,
2054    )
2055}
2056
2057fn history_mask_tensor<B: Backend>(
2058    history: &[[i64; HEARTMULA_PARALLEL_TOKENS]],
2059    device: &B::Device,
2060) -> Tensor<B, 3, Bool> {
2061    let flattened = history
2062        .iter()
2063        .flat_map(|row| {
2064            let has_audio_tokens = row[..HEARTMULA_AUDIO_CODEBOOKS]
2065                .iter()
2066                .any(|token| *token != 0);
2067            row.iter().enumerate().map(move |(index, token)| {
2068                if index < HEARTMULA_AUDIO_CODEBOOKS {
2069                    *token != 0
2070                } else if index == HEARTMULA_AUDIO_CODEBOOKS {
2071                    !has_audio_tokens
2072                } else {
2073                    false
2074                }
2075            })
2076        })
2077        .collect::<Vec<_>>();
2078    Tensor::<B, 3, Bool>::from_data(
2079        TensorData::new(flattened, [1, history.len(), HEARTMULA_PARALLEL_TOKENS]),
2080        device,
2081    )
2082}
2083
2084fn single_position_tensor<B: Backend>(position: i64, device: &B::Device) -> Tensor<B, 2, Int> {
2085    Tensor::<B, 2, Int>::from_data([[position]], device)
2086}
2087
2088fn position_tensor<B: Backend>(positions: Vec<i64>, device: &B::Device) -> Tensor<B, 2, Int> {
2089    let len = positions.len();
2090    Tensor::<B, 1, Int>::from_data(TensorData::new(positions, [len]), device).reshape([1, len])
2091}
2092
2093fn repeat_kv_heads<B: Backend>(tensor: Tensor<B, 4>, repeats: usize) -> Tensor<B, 4> {
2094    let [batch, seq_len, heads, head_dim] = tensor.dims();
2095    tensor
2096        .unsqueeze_dim::<5>(3)
2097        .repeat_dim(3, repeats)
2098        .reshape([batch, seq_len, heads * repeats, head_dim])
2099}
2100
2101fn repeat_cached_kv_heads<B: Backend>(tensor: Tensor<B, 4>, repeats: usize) -> Tensor<B, 4> {
2102    let [batch, heads, seq_len, head_dim] = tensor.dims();
2103    tensor
2104        .unsqueeze_dim::<5>(2)
2105        .repeat_dim(2, repeats)
2106        .reshape([batch, heads * repeats, seq_len, head_dim])
2107}
2108
2109fn apply_scaled_rope<B: Backend>(
2110    tensor: Tensor<B, 4>,
2111    positions: &Tensor<B, 2, Int>,
2112) -> Tensor<B, 4> {
2113    let [batch, seq_len, num_heads, head_dim] = tensor.dims();
2114    let pos = positions
2115        .clone()
2116        .to_data()
2117        .to_vec::<i64>()
2118        .expect("positions should be materializable");
2119    let cache = scaled_rope_cache::<B>(&tensor.device(), &pos, head_dim)
2120        .reshape([1, seq_len, 1, head_dim / 2, 2])
2121        .repeat_dim(0, batch);
2122    let reshaped = tensor.reshape([batch, seq_len, num_heads, head_dim / 2, 2]);
2123    Tensor::cat(
2124        vec![
2125            (reshaped
2126                .clone()
2127                .slice([0..batch, 0..seq_len, 0..num_heads, 0..head_dim / 2, 0..1])
2128                * cache
2129                    .clone()
2130                    .slice([0..batch, 0..seq_len, 0..1, 0..head_dim / 2, 0..1]))
2131                - (reshaped.clone().slice([
2132                    0..batch,
2133                    0..seq_len,
2134                    0..num_heads,
2135                    0..head_dim / 2,
2136                    1..2,
2137                ]) * cache
2138                    .clone()
2139                    .slice([0..batch, 0..seq_len, 0..1, 0..head_dim / 2, 1..2])),
2140            (reshaped
2141                .clone()
2142                .slice([0..batch, 0..seq_len, 0..num_heads, 0..head_dim / 2, 1..2])
2143                * cache
2144                    .clone()
2145                    .slice([0..batch, 0..seq_len, 0..1, 0..head_dim / 2, 0..1]))
2146                + (reshaped.slice([0..batch, 0..seq_len, 0..num_heads, 0..head_dim / 2, 0..1])
2147                    * cache.slice([0..batch, 0..seq_len, 0..1, 0..head_dim / 2, 1..2])),
2148        ],
2149        4,
2150    )
2151    .reshape([batch, seq_len, num_heads, head_dim])
2152}
2153
2154fn scaled_rope_cache<B: Backend>(
2155    device: &B::Device,
2156    positions: &[i64],
2157    head_dim: usize,
2158) -> Tensor<B, 3> {
2159    let theta = scaled_theta(head_dim);
2160    let mut values = Vec::with_capacity(positions.len() * (head_dim / 2) * 2);
2161    for &pos in positions {
2162        for &freq in &theta {
2163            let angle = pos as f32 * freq;
2164            values.push(angle.cos());
2165            values.push(angle.sin());
2166        }
2167    }
2168    Tensor::<B, 3>::from_data(
2169        TensorData::new(values, [positions.len(), head_dim / 2, 2]),
2170        device,
2171    )
2172}
2173
2174fn scaled_theta(head_dim: usize) -> Vec<f32> {
2175    (0..head_dim)
2176        .step_by(2)
2177        .map(|index| {
2178            let exponent = index as f32 / head_dim as f32;
2179            let freq = HEARTMULA_ROPE_BASE.powf(-exponent);
2180            let wavelength = 2.0 * std::f32::consts::PI / freq;
2181            let low_freq_wavelen = HEARTMULA_OLD_CONTEXT_LEN / HEARTMULA_LOW_FREQ_FACTOR;
2182            let high_freq_wavelen = HEARTMULA_OLD_CONTEXT_LEN / HEARTMULA_HIGH_FREQ_FACTOR;
2183            if wavelength < high_freq_wavelen {
2184                freq
2185            } else if wavelength > low_freq_wavelen {
2186                freq / HEARTMULA_ROPE_SCALE_FACTOR
2187            } else {
2188                let smooth = (HEARTMULA_OLD_CONTEXT_LEN / wavelength - HEARTMULA_LOW_FREQ_FACTOR)
2189                    / (HEARTMULA_HIGH_FREQ_FACTOR - HEARTMULA_LOW_FREQ_FACTOR);
2190                (1.0 - smooth) * freq / HEARTMULA_ROPE_SCALE_FACTOR + smooth * freq
2191            }
2192        })
2193        .collect()
2194}
2195
2196fn causal_mask<B: Backend>(seq_len: usize, device: &B::Device) -> Tensor<B, 4, Bool> {
2197    let mut mask = Vec::with_capacity(seq_len * seq_len);
2198    for row in 0..seq_len {
2199        for col in 0..seq_len {
2200            mask.push(col > row);
2201        }
2202    }
2203    Tensor::<B, 4, Bool>::from_data(TensorData::new(mask, [1, 1, seq_len, seq_len]), device)
2204}
2205
2206fn take_last_token<B: Backend>(hidden: Tensor<B, 3>) -> Tensor<B, 2> {
2207    let [batch, seq_len, hidden_size] = hidden.dims();
2208    hidden
2209        .slice([0..batch, seq_len - 1..seq_len, 0..hidden_size])
2210        .reshape([batch, hidden_size])
2211}
2212
2213fn tensor_to_f32_vec<B: Backend, const D: usize>(tensor: Tensor<B, D>) -> Result<Vec<f32>> {
2214    tensor
2215        .cast(DType::F32)
2216        .to_data()
2217        .to_vec::<f32>()
2218        .map_err(|e| anyhow!("failed to materialize tensor as f32: {:?}", e))
2219}
2220
2221fn sample_token<B: Backend>(logits: &Tensor<B, 2>, temperature: f32, topk: usize) -> Result<i64> {
2222    use burn::tensor::Distribution;
2223    use burn::tensor::activation::softmax;
2224
2225    if topk <= 1 {
2226        return argmax_token(logits);
2227    }
2228
2229    let scaled = logits.clone() / temperature;
2230
2231    let vocab_size = scaled.dims()[1];
2232    let k = topk.min(vocab_size).max(2);
2233
2234    let (topk_values, topk_indices) = scaled.clone().topk_with_indices(k, 1);
2235
2236    let probs = softmax(topk_values, 1);
2237
2238    let uniform = Tensor::<B, 2>::random([1, k], Distribution::Uniform(0.0, 1.0), &probs.device())
2239        .cast(burn::tensor::DType::F32);
2240    let uniform_data = uniform.to_data();
2241    let uniform_vec: Vec<f32> = uniform_data
2242        .to_vec()
2243        .map_err(|e| anyhow!("failed to get uniform random data: {:?}", e))?;
2244    let probs_data = probs.cast(burn::tensor::DType::F32).to_data();
2245    let probs_vec: Vec<f32> = probs_data
2246        .to_vec()
2247        .map_err(|e| anyhow!("failed to get probability data: {:?}", e))?;
2248
2249    let mut selected_idx = 0usize;
2250    let mut best_score = f32::NEG_INFINITY;
2251    for (i, (&u, &p)) in uniform_vec.iter().zip(probs_vec.iter()).enumerate() {
2252        let clamped = u.clamp(f32::MIN_POSITIVE, 1.0);
2253        let q = -clamped.ln();
2254        let score = p / q;
2255        if score > best_score {
2256            best_score = score;
2257            selected_idx = i;
2258        }
2259    }
2260
2261    let token_data = topk_indices
2262        .slice([0..1, selected_idx..selected_idx + 1])
2263        .to_data();
2264    let token_vec: Vec<i64> = token_data
2265        .to_vec()
2266        .map_err(|_| anyhow!("failed to get token"))?;
2267    let token = token_vec[0];
2268
2269    Ok(token)
2270}
2271
2272fn argmax_token<B: Backend>(logits: &Tensor<B, 2>) -> Result<i64> {
2273    let logits_data = logits.clone().cast(DType::F32).to_data();
2274    let logits_vec: Vec<f32> = logits_data
2275        .to_vec()
2276        .map_err(|e| anyhow!("failed to get logits for argmax: {:?}", e))?;
2277    let dims = logits.dims();
2278    let vocab_size = *dims
2279        .last()
2280        .ok_or_else(|| anyhow!("argmax_token expected non-empty logits shape"))?;
2281    if vocab_size == 0 || logits_vec.is_empty() {
2282        return Err(anyhow!("argmax_token received empty logits"));
2283    }
2284    let row = &logits_vec[..vocab_size];
2285    let mut best_index = 0usize;
2286    let mut best_value = f32::NEG_INFINITY;
2287    for (index, &value) in row.iter().enumerate() {
2288        if value > best_value {
2289            best_value = value;
2290            best_index = index;
2291        }
2292    }
2293    Ok(best_index as i64)
2294}
2295
2296fn linear_no_bias<B: Backend>(device: &B::Device, d_input: usize, d_output: usize) -> Linear<B> {
2297    LinearConfig::new(d_input, d_output)
2298        .with_bias(false)
2299        .with_layout(LinearLayout::Col)
2300        .init(device)
2301}
2302
2303fn linear_with_bias<B: Backend>(device: &B::Device, d_input: usize, d_output: usize) -> Linear<B> {
2304    LinearConfig::new(d_input, d_output)
2305        .with_layout(LinearLayout::Col)
2306        .init(device)
2307}
2308
2309fn uninitialized_param<B: Backend, const D: usize>(
2310    shape: [usize; D],
2311    device: &B::Device,
2312) -> Param<Tensor<B, D>> {
2313    Param::uninitialized(
2314        ParamId::new(),
2315        move |device, _require_grad| Tensor::<B, D>::zeros(shape, device),
2316        device.clone(),
2317        false,
2318        shape.into(),
2319    )
2320}
2321
2322#[cfg(test)]
2323mod tests {
2324    use super::*;
2325
2326    #[test]
2327    fn default_tags_returns_expected() {
2328        assert_eq!(super::default_tags(), "<tag></tag>");
2329    }
2330
2331    #[test]
2332    fn normalize_tags_adds_wrappers() {
2333        let input = "pop, electronic";
2334        let result = super::normalize_tags(input);
2335        assert_eq!(result, "<tag>pop,electronic</tag>");
2336    }
2337
2338    #[test]
2339    fn normalize_tags_preserves_existing_wrappers() {
2340        let input = "<tag>pop</tag>";
2341        let result = super::normalize_tags(input);
2342        assert_eq!(result, "<tag>pop</tag>");
2343    }
2344
2345    #[test]
2346    fn normalize_tags_removes_spaces_after_commas() {
2347        let input = "pop, rock, jazz";
2348        let result = super::normalize_tags(input);
2349        assert_eq!(result, "<tag>pop,rock,jazz</tag>");
2350    }
2351
2352    #[test]
2353    fn normalize_tags_converts_to_lowercase() {
2354        let input = "POP, ROCK";
2355        let result = super::normalize_tags(input);
2356        assert_eq!(result, "<tag>pop,rock</tag>");
2357    }
2358
2359    #[test]
2360    fn normalize_tags_trims_whitespace() {
2361        let input = "  pop, rock  ";
2362        let result = super::normalize_tags(input);
2363        assert_eq!(result, "<tag>pop,rock</tag>");
2364    }
2365
2366    #[test]
2367    fn build_prompt_history_basic() {
2368        let text_bos_id = 1_i64;
2369        let text_eos_id = 2_i64;
2370        let lyrics_ids = vec![10, 11, 12];
2371        let tags_ids = vec![20, 21];
2372
2373        let history = super::build_prompt_history(text_bos_id, text_eos_id, &lyrics_ids, &tags_ids);
2374
2375        assert_eq!(history.len(), 10);
2376
2377        assert_eq!(history[0][HEARTMULA_AUDIO_CODEBOOKS], text_bos_id);
2378
2379        assert!(history[4].iter().all(|&x| x == 0));
2380
2381        assert_eq!(history[5][HEARTMULA_AUDIO_CODEBOOKS], text_bos_id);
2382    }
2383
2384    #[test]
2385    fn normalize_text_ids_adds_bos_and_eos() {
2386        let text_bos_id = 1_i64;
2387        let text_eos_id = 2_i64;
2388        let ids = vec![10, 11, 12];
2389
2390        let result = super::normalize_text_ids(text_bos_id, text_eos_id, &ids);
2391
2392        assert_eq!(result[0], text_bos_id);
2393        assert_eq!(result[result.len() - 1], text_eos_id);
2394        assert_eq!(result, vec![1, 10, 11, 12, 2]);
2395    }
2396
2397    #[test]
2398    fn normalize_text_ids_preserves_existing_bos_eos() {
2399        let text_bos_id = 1_i64;
2400        let text_eos_id = 2_i64;
2401        let ids = vec![1, 10, 11, 12, 2];
2402
2403        let result = super::normalize_text_ids(text_bos_id, text_eos_id, &ids);
2404
2405        assert_eq!(result, vec![1, 10, 11, 12, 2]);
2406    }
2407
2408    #[test]
2409    fn build_audio_history_row() {
2410        let frame = vec![100, 200, 300, 400, 500, 600, 700, 800];
2411        let empty_id = 0_i64;
2412
2413        let row = super::build_audio_history_row(&frame, empty_id);
2414
2415        for i in 0..HEARTMULA_AUDIO_CODEBOOKS {
2416            assert_eq!(row[i], frame[i] as i64);
2417        }
2418
2419        assert_eq!(row[HEARTMULA_AUDIO_CODEBOOKS], empty_id);
2420    }
2421
2422    #[test]
2423    fn build_audio_history_row_with_short_frame() {
2424        let frame = vec![100, 200];
2425        let empty_id = 999_i64;
2426
2427        let row = super::build_audio_history_row(&frame, empty_id);
2428
2429        assert_eq!(row[0], 100);
2430        assert_eq!(row[1], 200);
2431
2432        for item in row.iter().take(HEARTMULA_AUDIO_CODEBOOKS).skip(2) {
2433            assert_eq!(*item, empty_id);
2434        }
2435    }
2436
2437    #[test]
2438    fn single_position_tensor() {
2439        use burn::backend::ndarray::NdArray;
2440
2441        let device = burn::prelude::Device::<NdArray<f32>>::default();
2442        let tensor = super::single_position_tensor::<NdArray<f32>>(42, &device);
2443
2444        assert_eq!(tensor.dims(), [1, 1]);
2445        let data = tensor.to_data().to_vec::<i64>().unwrap();
2446        assert_eq!(data[0], 42);
2447    }
2448
2449    #[test]
2450    fn position_tensor() {
2451        use burn::backend::ndarray::NdArray;
2452
2453        let device = burn::prelude::Device::<NdArray<f32>>::default();
2454        let positions = vec![0, 1, 2, 3, 4];
2455        let tensor = super::position_tensor::<NdArray<f32>>(positions, &device);
2456
2457        assert_eq!(tensor.dims(), [1, 5]);
2458        let data = tensor.to_data().to_vec::<i64>().unwrap();
2459        assert_eq!(data, vec![0, 1, 2, 3, 4]);
2460    }
2461
2462    #[test]
2463    fn repeat_kv_heads() {
2464        use burn::backend::ndarray::NdArray;
2465
2466        let device = burn::prelude::Device::<NdArray<f32>>::default();
2467        let tensor = Tensor::<NdArray<f32>, 4>::from_data(
2468            TensorData::new(vec![1.0; 32], [1, 2, 4, 4]),
2469            &device,
2470        );
2471
2472        let repeated = super::repeat_kv_heads(tensor, 2);
2473        assert_eq!(repeated.dims(), [1, 2, 8, 4]);
2474    }
2475
2476    #[test]
2477    fn repeat_cached_kv_heads() {
2478        use burn::backend::ndarray::NdArray;
2479
2480        let device = burn::prelude::Device::<NdArray<f32>>::default();
2481        let tensor = Tensor::<NdArray<f32>, 4>::from_data(
2482            TensorData::new(vec![1.0; 32], [1, 4, 2, 4]),
2483            &device,
2484        );
2485
2486        let repeated = super::repeat_cached_kv_heads(tensor, 3);
2487        assert_eq!(repeated.dims(), [1, 12, 2, 4]);
2488    }
2489
2490    #[test]
2491    fn take_last_token() {
2492        use burn::backend::ndarray::NdArray;
2493
2494        let device = burn::prelude::Device::<NdArray<f32>>::default();
2495        let tensor = Tensor::<NdArray<f32>, 3>::from_data(
2496            TensorData::new(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], [1, 3, 2]),
2497            &device,
2498        );
2499
2500        let last = super::take_last_token(tensor);
2501        assert_eq!(last.dims(), [1, 2]);
2502    }
2503
2504    #[test]
2505    fn tensor_to_f32_vec_success() {
2506        use burn::backend::ndarray::NdArray;
2507
2508        let device = burn::prelude::Device::<NdArray<f32>>::default();
2509        let tensor = Tensor::<NdArray<f32>, 2>::from_data(
2510            TensorData::new(vec![1.0, 2.0, 3.0, 4.0], [2, 2]),
2511            &device,
2512        );
2513
2514        let vec = super::tensor_to_f32_vec(tensor).unwrap();
2515        assert_eq!(vec, vec![1.0, 2.0, 3.0, 4.0]);
2516    }
2517
2518    #[test]
2519    fn write_frames_json_creates_valid_json() {
2520        use std::io::Read;
2521
2522        let temp_dir = std::env::temp_dir();
2523        let path = temp_dir.join("test_frames.json");
2524
2525        let frames: Vec<Vec<i64>> = vec![
2526            vec![1, 2, 3, 4, 5, 6, 7, 8],
2527            vec![9, 10, 11, 12, 13, 14, 15, 16],
2528        ];
2529
2530        super::write_frames_json(&path, "test lyrics", "test tags", &frames).unwrap();
2531
2532        let mut file = std::fs::File::open(&path).unwrap();
2533        let mut contents = String::new();
2534        file.read_to_string(&mut contents).unwrap();
2535
2536        assert!(contents.contains("heartmula"));
2537        assert!(contents.contains("test lyrics"));
2538        assert!(contents.contains("test tags"));
2539        assert!(contents.contains("frame_count"));
2540        assert!(contents.contains("48000"));
2541
2542        std::fs::remove_file(&path).unwrap();
2543    }
2544
2545    #[test]
2546    fn heartmula_generation_config_defaults() {
2547        let lyrics_ids: &[i64] = &[1, 2, 3];
2548        let tags_ids: &[i64] = &[4, 5];
2549
2550        let config = HeartmulaGenerationConfig {
2551            text_bos_id: 1,
2552            text_eos_id: 2,
2553            audio_eos_id: 1000,
2554            empty_id: 0,
2555            lyrics_ids,
2556            tags_ids,
2557            max_audio_frames: 100,
2558            temperature: 0.8,
2559            topk: 25,
2560            cfg_scale: 2.0,
2561            progress_callback: None,
2562        };
2563
2564        assert_eq!(config.temperature, 0.8);
2565        assert_eq!(config.topk, 25);
2566        assert_eq!(config.cfg_scale, 2.0);
2567        assert_eq!(config.max_audio_frames, 100);
2568    }
2569
2570    #[test]
2571    fn splitmix64_produces_deterministic_sequence() {
2572        let mut state1 = 123456789_u64;
2573        let mut state2 = 123456789_u64;
2574
2575        for _ in 0..10 {
2576            assert_eq!(
2577                super::splitmix64_next(&mut state1),
2578                super::splitmix64_next(&mut state2)
2579            );
2580        }
2581    }
2582
2583    #[test]
2584    fn uniform01_open_produces_values_in_range() {
2585        let mut state = 123456789_u64;
2586
2587        for _ in 0..100 {
2588            let value = super::uniform01_open(&mut state);
2589            assert!(value > 0.0);
2590            assert!(value < 1.0);
2591        }
2592    }
2593
2594    #[test]
2595    fn generate_decoder_latent_data_deterministic() {
2596        let data1 = super::generate_decoder_latent_data(42, 100);
2597        let data2 = super::generate_decoder_latent_data(42, 100);
2598
2599        assert_eq!(data1, data2);
2600        assert_eq!(data1.len(), 100);
2601    }
2602
2603    #[test]
2604    fn scaled_theta_produces_expected_length() {
2605        let head_dim = 64;
2606        let theta = super::scaled_theta(head_dim);
2607
2608        assert_eq!(theta.len(), head_dim / 2);
2609    }
2610
2611    #[test]
2612    fn scaled_theta_frequency_scaling() {
2613        let theta_64 = super::scaled_theta(64);
2614        let theta_128 = super::scaled_theta(128);
2615
2616        for i in 1..theta_64.len() {
2617            assert!(theta_64[i] <= theta_64[i - 1]);
2618        }
2619
2620        for i in 1..theta_128.len() {
2621            assert!(theta_128[i] <= theta_128[i - 1]);
2622        }
2623    }
2624
2625    #[test]
2626    fn heartmula_rms_norm_forward() {
2627        use burn::backend::ndarray::NdArray;
2628
2629        let device = burn::prelude::Device::<NdArray<f32>>::default();
2630        let norm = super::HeartmulaRmsNorm::<NdArray<f32>>::new(&device, 64, 1e-5);
2631
2632        let input = Tensor::<NdArray<f32>, 3>::ones([1, 4, 64], &device);
2633        let output = norm.forward(input);
2634
2635        assert_eq!(output.dims(), [1, 4, 64]);
2636    }
2637
2638    #[test]
2639    fn heartmula_mlp_forward() {
2640        use burn::backend::ndarray::NdArray;
2641
2642        let device = burn::prelude::Device::<NdArray<f32>>::default();
2643        let mlp = super::HeartmulaMlp::<NdArray<f32>>::new(&device);
2644
2645        let input = Tensor::<NdArray<f32>, 3>::ones([1, 4, HEARTMULA_HIDDEN_SIZE], &device);
2646        let output = mlp.forward(input);
2647
2648        assert_eq!(output.dims(), [1, 4, HEARTMULA_HIDDEN_SIZE]);
2649    }
2650
2651    #[test]
2652    fn heartmula_attention_new() {
2653        use burn::backend::ndarray::NdArray;
2654
2655        let device = burn::prelude::Device::<NdArray<f32>>::default();
2656        let attn = super::HeartmulaAttention::<NdArray<f32>>::new(
2657            &device,
2658            HEARTMULA_BACKBONE_HEADS,
2659            HEARTMULA_BACKBONE_KV_HEADS,
2660        );
2661
2662        assert_eq!(attn.meta.num_heads, HEARTMULA_BACKBONE_HEADS);
2663        assert_eq!(attn.meta.num_kv_heads, HEARTMULA_BACKBONE_KV_HEADS);
2664    }
2665
2666    #[test]
2667    fn heartmula_transformer_layer_new() {
2668        use burn::backend::ndarray::NdArray;
2669
2670        let device = burn::prelude::Device::<NdArray<f32>>::default();
2671        let layer = super::HeartmulaTransformerLayer::<NdArray<f32>>::new(
2672            &device,
2673            HEARTMULA_BACKBONE_HEADS,
2674            HEARTMULA_BACKBONE_KV_HEADS,
2675        );
2676
2677        assert_eq!(layer.attn.meta.num_heads, HEARTMULA_BACKBONE_HEADS);
2678    }
2679
2680    #[test]
2681    fn heartmula_transformer_new() {
2682        use burn::backend::ndarray::NdArray;
2683
2684        let device = burn::prelude::Device::<NdArray<f32>>::default();
2685        let transformer = super::HeartmulaTransformer::<NdArray<f32>>::new(
2686            &device,
2687            HEARTMULA_BACKBONE_LAYERS,
2688            HEARTMULA_BACKBONE_HEADS,
2689            HEARTMULA_BACKBONE_KV_HEADS,
2690        );
2691
2692        assert_eq!(transformer.layers.len(), HEARTMULA_BACKBONE_LAYERS);
2693    }
2694
2695    #[test]
2696    fn heartmula_model_new() {
2697        use burn::backend::ndarray::NdArray;
2698
2699        let device = burn::prelude::Device::<NdArray<f32>>::default();
2700        let model = super::HeartmulaModel::<NdArray<f32>>::new(&device, 1000, 1024);
2701
2702        assert_eq!(model.audio_head.dims()[0], HEARTMULA_AUDIO_CODEBOOKS - 1);
2703        assert_eq!(model.audio_head.dims()[1], HEARTMULA_HIDDEN_SIZE);
2704        assert_eq!(model.audio_head.dims()[2], 1024);
2705    }
2706
2707    #[test]
2708    fn heartmula_model_audio_vocab_size() {
2709        use burn::backend::ndarray::NdArray;
2710
2711        let device = burn::prelude::Device::<NdArray<f32>>::default();
2712        let model = super::HeartmulaModel::<NdArray<f32>>::new(&device, 1000, 1024);
2713
2714        assert_eq!(model.audio_vocab_size(), 1024);
2715    }
2716
2717    #[test]
2718    fn history_tokens_tensor_shape() {
2719        use burn::backend::ndarray::NdArray;
2720
2721        let device = burn::prelude::Device::<NdArray<f32>>::default();
2722        let history: Vec<[i64; HEARTMULA_PARALLEL_TOKENS]> = vec![
2723            [1, 2, 3, 4, 5, 6, 7, 8, 100],
2724            [9, 10, 11, 12, 13, 14, 15, 16, 101],
2725        ];
2726
2727        let tensor = super::history_tokens_tensor::<NdArray<f32>>(&history, &device);
2728        assert_eq!(tensor.dims(), [1, 2, HEARTMULA_PARALLEL_TOKENS]);
2729    }
2730
2731    #[test]
2732    fn history_mask_tensor_shape() {
2733        use burn::backend::ndarray::NdArray;
2734
2735        let device = burn::prelude::Device::<NdArray<f32>>::default();
2736        let history: Vec<[i64; HEARTMULA_PARALLEL_TOKENS]> =
2737            vec![[1, 2, 3, 4, 5, 6, 7, 8, 100], [0, 0, 0, 0, 0, 0, 0, 0, 101]];
2738
2739        let tensor = super::history_mask_tensor::<NdArray<f32>>(&history, &device);
2740        assert_eq!(tensor.dims(), [1, 2, HEARTMULA_PARALLEL_TOKENS]);
2741    }
2742
2743    #[test]
2744    fn causal_mask_shape() {
2745        use burn::backend::ndarray::NdArray;
2746
2747        let device = burn::prelude::Device::<NdArray<f32>>::default();
2748        let mask = super::causal_mask::<NdArray<f32>>(5, &device);
2749
2750        assert_eq!(mask.dims(), [1, 1, 5, 5]);
2751    }
2752
2753    #[test]
2754    fn causal_mask_values() {
2755        use burn::backend::ndarray::NdArray;
2756
2757        let device = burn::prelude::Device::<NdArray<f32>>::default();
2758        let mask = super::causal_mask::<NdArray<f32>>(3, &device);
2759        let data = mask.to_data().to_vec::<bool>().unwrap();
2760
2761        assert!(!data[0]);
2762        assert!(data[1]);
2763        assert!(data[2]);
2764        assert!(!data[3]);
2765        assert!(!data[4]);
2766        assert!(data[5]);
2767        assert!(!data[6]);
2768        assert!(!data[7]);
2769        assert!(!data[8]);
2770    }
2771
2772    #[test]
2773    fn apply_scaled_rope_preserves_shape() {
2774        use burn::backend::ndarray::NdArray;
2775
2776        let device = burn::prelude::Device::<NdArray<f32>>::default();
2777        let tensor = Tensor::<NdArray<f32>, 4>::ones([1, 4, 8, 64], &device);
2778        let positions = Tensor::<NdArray<f32>, 2, Int>::from_data(
2779            TensorData::new(vec![0, 1, 2, 3], [1, 4]),
2780            &device,
2781        );
2782
2783        let rotated = super::apply_scaled_rope(tensor, &positions);
2784        assert_eq!(rotated.dims(), [1, 4, 8, 64]);
2785    }
2786
2787    #[test]
2788    fn scaled_rope_cache_shape() {
2789        use burn::backend::ndarray::NdArray;
2790
2791        let device = burn::prelude::Device::<NdArray<f32>>::default();
2792        let positions = vec![0, 1, 2, 3, 4];
2793        let cache = super::scaled_rope_cache::<NdArray<f32>>(&device, &positions, 64);
2794
2795        assert_eq!(cache.dims(), [5, 32, 2]);
2796    }
2797
2798    #[test]
2799    fn splice_sequence_token_middle() {
2800        use burn::backend::ndarray::NdArray;
2801
2802        let device = burn::prelude::Device::<NdArray<f32>>::default();
2803        let hidden = Tensor::<NdArray<f32>, 3>::from_data(
2804            TensorData::new(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], [1, 3, 2]),
2805            &device,
2806        );
2807        let replacement = Tensor::<NdArray<f32>, 3>::from_data(
2808            TensorData::new(vec![9.0, 9.0], [1, 1, 2]),
2809            &device,
2810        );
2811
2812        let result = super::splice_sequence_token(hidden, replacement, 1);
2813        assert_eq!(result.dims(), [1, 3, 2]);
2814    }
2815
2816    #[test]
2817    fn splice_sequence_token_at_start() {
2818        use burn::backend::ndarray::NdArray;
2819
2820        let device = burn::prelude::Device::<NdArray<f32>>::default();
2821        let hidden = Tensor::<NdArray<f32>, 3>::from_data(
2822            TensorData::new(vec![1.0, 2.0, 3.0, 4.0], [1, 2, 2]),
2823            &device,
2824        );
2825        let replacement = Tensor::<NdArray<f32>, 3>::from_data(
2826            TensorData::new(vec![9.0, 9.0], [1, 1, 2]),
2827            &device,
2828        );
2829
2830        let result = super::splice_sequence_token(hidden, replacement, 0);
2831        assert_eq!(result.dims(), [1, 2, 2]);
2832    }
2833
2834    #[test]
2835    fn splice_sequence_token_at_end() {
2836        use burn::backend::ndarray::NdArray;
2837
2838        let device = burn::prelude::Device::<NdArray<f32>>::default();
2839        let hidden = Tensor::<NdArray<f32>, 3>::from_data(
2840            TensorData::new(vec![1.0, 2.0, 3.0, 4.0], [1, 2, 2]),
2841            &device,
2842        );
2843        let replacement = Tensor::<NdArray<f32>, 3>::from_data(
2844            TensorData::new(vec![9.0, 9.0], [1, 1, 2]),
2845            &device,
2846        );
2847
2848        let result = super::splice_sequence_token(hidden, replacement, 1);
2849        assert_eq!(result.dims(), [1, 2, 2]);
2850    }
2851}