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}