use std::path::{Path, PathBuf};
use anyhow::{Context, Result, bail};
use burn::module::{Module, Param};
use burn::prelude::Backend;
use burn::tensor::{Int, Tensor, TensorData};
use burn_store::{BurnpackStore, ModuleStore};
use super::condition::AceStepCondition;
use super::config::AceStepConfig;
use super::dit::{AceStepDiT, SFT_TIMESTEPS, TURBO_TIMESTEPS};
use super::lm::{self, AceStepLm, AudioCodeVocab, SamplingConfig};
use super::qwen3::{Qwen3Config, Qwen3Model};
use super::vae::{OobleckDecoder, OobleckVaeConfig};
const TIMBRE_REFERENCE_FRAMES: usize = 750;
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
pub enum AceStepVariant {
#[default]
Turbo,
Sft,
}
impl AceStepVariant {
fn prefix(self) -> &'static str {
match self {
Self::Turbo => "",
Self::Sft => "sft-",
}
}
}
#[derive(Clone, Debug)]
pub struct AceStepModelPaths {
pub text_encoder_bpk: PathBuf,
pub text_encoder_config: PathBuf,
pub lm_bpk: PathBuf,
pub lm_config: PathBuf,
pub lm_tokenizer: PathBuf,
pub dit_bpk: PathBuf,
pub dit_config: PathBuf,
pub condition_bpk: PathBuf,
pub vae_bpk: PathBuf,
pub vae_config: PathBuf,
pub silence_latent_bpk: PathBuf,
pub tokenizer_json: PathBuf,
}
impl AceStepModelPaths {
pub fn required_relative_files(variant: AceStepVariant) -> &'static [&'static str] {
const TURBO: [&str; 12] = [
"qwen3-encoder.bpk",
"qwen3_config.json",
"tokenizer.json",
"acestep-vae.bpk",
"vae_config.json",
"silence_latent.bpk",
"acestep-lm.bpk",
"lm_config.json",
"lm_tokenizer.json",
"acestep-dit.bpk",
"dit_config.json",
"acestep-condition.bpk",
];
const SFT: [&str; 12] = [
"qwen3-encoder.bpk",
"qwen3_config.json",
"tokenizer.json",
"acestep-vae.bpk",
"vae_config.json",
"silence_latent.bpk",
"sft-acestep-lm.bpk",
"sft-lm_config.json",
"sft-lm_tokenizer.json",
"sft-acestep-dit.bpk",
"sft-dit_config.json",
"sft-acestep-condition.bpk",
];
match variant {
AceStepVariant::Turbo => &TURBO,
AceStepVariant::Sft => &SFT,
}
}
pub fn resolve(model_dir: &Path, variant: AceStepVariant) -> Result<Self> {
let prefix = variant.prefix();
let paths = Self {
text_encoder_bpk: model_dir.join("qwen3-encoder.bpk"),
text_encoder_config: model_dir.join("qwen3_config.json"),
lm_bpk: model_dir.join(format!("{prefix}acestep-lm.bpk")),
lm_config: model_dir.join(format!("{prefix}lm_config.json")),
lm_tokenizer: model_dir.join(format!("{prefix}lm_tokenizer.json")),
dit_bpk: model_dir.join(format!("{prefix}acestep-dit.bpk")),
dit_config: model_dir.join(format!("{prefix}dit_config.json")),
condition_bpk: model_dir.join(format!("{prefix}acestep-condition.bpk")),
vae_bpk: model_dir.join("acestep-vae.bpk"),
vae_config: model_dir.join("vae_config.json"),
silence_latent_bpk: model_dir.join("silence_latent.bpk"),
tokenizer_json: model_dir.join("tokenizer.json"),
};
let relative = Self::required_relative_files(variant);
let missing: Vec<&'static str> = relative
.iter()
.copied()
.zip(paths.all_absolute())
.filter(|(_, absolute)| !absolute.exists())
.map(|(relative, _)| relative)
.collect();
if !missing.is_empty() {
bail!(
"ACE-Step model directory {} is missing required files: {}",
model_dir.display(),
missing.join(", ")
);
}
Ok(paths)
}
fn all_absolute(&self) -> [&PathBuf; 12] {
[
&self.text_encoder_bpk,
&self.text_encoder_config,
&self.lm_bpk,
&self.lm_config,
&self.lm_tokenizer,
&self.dit_bpk,
&self.dit_config,
&self.condition_bpk,
&self.vae_bpk,
&self.vae_config,
&self.silence_latent_bpk,
&self.tokenizer_json,
]
}
}
#[derive(Module, Debug)]
pub struct SilenceLatent<B: Backend> {
pub silence_latent: Param<Tensor<B, 3>>,
}
impl<B: Backend> SilenceLatent<B> {
pub fn from_burnpack(path: &Path, device: &B::Device) -> Result<Self> {
let data = BurnpackStore::from_file(path)
.zero_copy(true)
.get_all_snapshots()
.with_context(|| format!("failed to read snapshots from {}", path.display()))?
.iter()
.find_map(|(_, snap)| {
(snap.full_path() == "silence_latent").then(|| snap.to_data().ok())
})
.flatten()
.ok_or_else(|| anyhow::anyhow!("missing silence_latent tensor in {}", path.display()))?
.convert::<f32>();
let tensor = Tensor::<B, 3>::from_data(data, device);
Ok(Self {
silence_latent: Param::from_tensor(tensor),
})
}
pub fn slice(&self, frames: usize) -> Tensor<B, 3> {
tile_frames(self.silence_latent.val(), frames)
}
pub fn timbre_reference(&self) -> Tensor<B, 3> {
self.slice(TIMBRE_REFERENCE_FRAMES)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct GenerateAudioMeta {
pub channels: usize,
pub frames: usize,
pub sample_rate_hz: u32,
pub steps: usize,
pub prompt_tokens: usize,
}
#[derive(Clone, Copy, Debug, Default, PartialEq)]
pub struct GenerateMetadata<'a> {
pub bpm: Option<f32>,
pub key_scale: Option<&'a str>,
pub time_signature: Option<&'a str>,
}
#[derive(Clone, Copy, Debug)]
pub struct DitStepStat {
pub t: f32,
pub xt_rms: f32,
pub v_rms: f32,
}
#[derive(Clone, Debug, Default)]
pub struct AceStepTrace {
pub text_prompt: String,
pub lyric_prompt: String,
pub cot_block: String,
pub lm_prompt: String,
pub codes: Vec<u32>,
pub enc_mean: f32,
pub enc_std: f32,
pub hints_latent_rms: f32,
pub hints_audio: Option<(Vec<f32>, usize, usize)>,
pub dit_steps: Vec<DitStepStat>,
pub final_latent_rms: f32,
}
pub struct AceStepPipeline<B: Backend> {
pub paths: AceStepModelPaths,
pub device: B::Device,
pub text_encoder: Qwen3Model<B>,
pub condition: AceStepCondition<B>,
pub dit: AceStepDiT<B>,
pub vae: OobleckDecoder<B>,
pub silence: SilenceLatent<B>,
pub dit_config: AceStepConfig,
pub vae_config: OobleckVaeConfig,
pub audio_code_vocab: AudioCodeVocab,
pub lm_f16: bool,
text_tokenizer: tokie::Tokenizer,
}
impl<B: Backend> AceStepPipeline<B> {
pub fn load(
paths: &AceStepModelPaths,
device: &B::Device,
progress: &mut dyn FnMut(&str, f32, &str),
) -> Result<Self> {
progress("loading", 0.0, "text encoder config");
let text_config = Qwen3Config::load(&paths.text_encoder_config)?;
progress("loading", 0.1, "text encoder weights");
let text_encoder =
Qwen3Model::from_burnpack(&text_config, &paths.text_encoder_bpk, device)?;
progress("loading", 0.4, "dit config");
let dit_config = AceStepConfig::load(&paths.dit_config)?;
progress("loading", 0.5, "condition stack weights");
let condition = AceStepCondition::from_burnpack(&dit_config, &paths.condition_bpk, device)?;
progress("loading", 0.7, "dit weights");
let dit = AceStepDiT::from_burnpack(&dit_config, &paths.dit_bpk, device)?;
progress("loading", 0.85, "vae config");
let vae_config = OobleckVaeConfig::load(&paths.vae_config)?;
progress("loading", 0.9, "vae decoder weights");
let vae = OobleckDecoder::from_burnpack(&vae_config, &paths.vae_bpk, device)?;
progress("loading", 0.95, "silence latent");
let silence = SilenceLatent::from_burnpack(&paths.silence_latent_bpk, device)?;
progress("loading", 0.98, "tokenizers");
let text_tokenizer = tokie::Tokenizer::from_json(&paths.tokenizer_json).map_err(|e| {
anyhow::anyhow!(
"failed to load tokenizer from {}: {e}",
paths.tokenizer_json.display()
)
})?;
let audio_code_vocab = AudioCodeVocab::from_tokenizer_json(&paths.lm_tokenizer)?;
progress("loading", 1.0, "ready");
Ok(Self {
paths: paths.clone(),
device: device.clone(),
text_encoder,
condition,
dit,
vae,
silence,
dit_config,
vae_config,
audio_code_vocab,
lm_f16: false,
text_tokenizer,
})
}
pub fn plan_codes(
&self,
caption: &str,
metadata: &GenerateMetadata<'_>,
duration_ms: usize,
seed: u64,
progress: &mut dyn FnMut(&str, f32, &str),
) -> Result<(Vec<u32>, String, String)> {
let duration_s = duration_ms / 1000;
let cot = lm::build_cot_block(
caption,
metadata.bpm,
metadata.key_scale,
metadata.time_signature,
duration_s,
);
let prompt = lm::build_codes_prompt(caption, &cot);
let prompt_ids = lm::tokenize_prompt(&self.paths.lm_tokenizer, &prompt)?;
let sampling = SamplingConfig::new(lm::code_count_for_duration(duration_s) * 6, seed);
let uncond_prompt_ids = if sampling.cfg_scale > 1.0 {
Some(lm::tokenize_prompt(
&self.paths.lm_tokenizer,
&lm::build_uncond_codes_prompt(),
)?)
} else {
None
};
progress("loading", 0.0, "lm planner config");
let lm_config = Qwen3Config::load(&self.paths.lm_config)?;
let use_f16 = self.lm_f16
&& B::name(&self.device).contains("wgpu")
&& std::env::var_os("MAOLAN_ACESTEP_LM_F32").is_none();
let codes = if use_f16 {
progress("loading", 0.05, "lm planner weights (f16)");
let device = Default::default();
let lm =
AceStepLm::<burn::backend::Wgpu<burn::tensor::f16, i64, u32>>::from_burnpack_cast(
&lm_config,
&self.paths.lm_bpk,
&device,
)?;
let codes = {
let mut lm_progress = |done: usize, total: usize| {
let fraction = if total == 0 {
1.0
} else {
done as f32 / total as f32
};
progress("lm", fraction, "planning audio codes");
};
lm.generate_codes(
&prompt_ids,
uncond_prompt_ids.as_deref(),
&self.audio_code_vocab,
&sampling,
Some(&mut lm_progress),
)
};
drop(lm);
burn::backend::Wgpu::<burn::tensor::f16, i64, u32>::memory_cleanup(&device);
codes
} else {
progress("loading", 0.05, "lm planner weights");
let lm = AceStepLm::<B>::from_burnpack(&lm_config, &self.paths.lm_bpk, &self.device)?;
let codes = {
let mut lm_progress = |done: usize, total: usize| {
let fraction = if total == 0 {
1.0
} else {
done as f32 / total as f32
};
progress("lm", fraction, "planning audio codes");
};
lm.generate_codes(
&prompt_ids,
uncond_prompt_ids.as_deref(),
&self.audio_code_vocab,
&sampling,
Some(&mut lm_progress),
)
};
drop(lm);
B::memory_cleanup(&self.device);
codes
};
Ok((codes, cot, prompt))
}
pub fn generate(
&self,
caption: &str,
metadata: &GenerateMetadata<'_>,
duration_ms: usize,
seed: u64,
progress: &mut dyn FnMut(&str, f32, &str),
) -> Result<(Tensor<B, 3>, GenerateAudioMeta)> {
self.generate_impl(caption, metadata, duration_ms, seed, progress, None)
}
pub fn generate_traced(
&self,
caption: &str,
metadata: &GenerateMetadata<'_>,
duration_ms: usize,
seed: u64,
progress: &mut dyn FnMut(&str, f32, &str),
trace: &mut AceStepTrace,
) -> Result<(Tensor<B, 3>, GenerateAudioMeta)> {
self.generate_impl(caption, metadata, duration_ms, seed, progress, Some(trace))
}
fn generate_impl(
&self,
caption: &str,
metadata: &GenerateMetadata<'_>,
duration_ms: usize,
seed: u64,
progress: &mut dyn FnMut(&str, f32, &str),
mut trace: Option<&mut AceStepTrace>,
) -> Result<(Tensor<B, 3>, GenerateAudioMeta)> {
let acoustic_dim = self.dit_config.audio_acoustic_hidden_dim;
let latent_frames = latent_frames_for_duration(duration_ms);
let duration_s = duration_ms / 1000;
let device = self.silence.silence_latent.device();
let metas_block = build_metas_block(
metadata.bpm,
metadata.key_scale,
metadata.time_signature,
duration_s,
);
progress("text-encoder", 0.0, "tokenizing caption");
let text_prompt = build_dit_text_prompt(caption, &metas_block);
if let Some(trace) = trace.as_deref_mut() {
trace.text_prompt = text_prompt.clone();
trace.lyric_prompt = INSTRUMENTAL_LYRIC_PROMPT.to_string();
}
let mut ids = self.text_tokenizer.encode(&text_prompt, false).ids;
ids.truncate(MAX_TEXT_TOKENS);
ids.push(lm::ENDOFTEXT_ID);
let prompt_tokens = ids.len();
let ids: Vec<i64> = ids.into_iter().map(i64::from).collect();
let ids_tensor =
Tensor::<B, 2, Int>::from_data(TensorData::new(ids, [1, prompt_tokens]), &device);
progress("text-encoder", 0.3, "encoding caption");
let text_hidden = self.text_encoder.forward(ids_tensor, true);
let mut lyric_ids = self
.text_tokenizer
.encode(INSTRUMENTAL_LYRIC_PROMPT, false)
.ids;
lyric_ids.truncate(MAX_LYRIC_TOKENS);
lyric_ids.push(lm::ENDOFTEXT_ID);
let lyric_len = lyric_ids.len();
let lyric_ids: Vec<i64> = lyric_ids.into_iter().map(i64::from).collect();
let lyric_ids_tensor =
Tensor::<B, 2, Int>::from_data(TensorData::new(lyric_ids, [1, lyric_len]), &device);
let lyric_hidden = self.text_encoder.embed_tokens.forward(lyric_ids_tensor);
let lyric_mask = Tensor::<B, 2, Int>::ones([1, lyric_len], &device);
progress("text-encoder", 1.0, "caption encoded");
progress("condition", 0.0, "encoding conditions");
let enc = self.condition.encode(
text_hidden,
lyric_hidden,
lyric_mask,
self.silence.timbre_reference(),
);
progress("condition", 1.0, "conditions encoded");
if let Some(trace) = trace.as_deref_mut() {
let values: Vec<f32> = enc
.clone()
.into_data()
.convert::<f32>()
.to_vec()
.map_err(|e| anyhow::anyhow!("failed to read encoder states: {e}"))?;
let mean = values.iter().sum::<f32>() / values.len() as f32;
let var =
values.iter().map(|v| (v - mean) * (v - mean)).sum::<f32>() / values.len() as f32;
trace.enc_mean = mean;
trace.enc_std = var.sqrt();
}
let codes = if let Some(raw) = std::env::var_os("MAOLAN_ACESTEP_CODES") {
raw.to_string_lossy()
.split(',')
.filter_map(|part| part.trim().parse::<u32>().ok())
.collect::<Vec<u32>>()
} else {
let (codes, cot, prompt) =
self.plan_codes(caption, metadata, duration_ms, seed, &mut *progress)?;
if let Some(trace) = trace.as_deref_mut() {
trace.cot_block = cot;
trace.lm_prompt = prompt;
trace.codes = codes.clone();
}
codes
};
let expected = lm::code_count_for_duration(duration_s).max(1);
let codes = pad_or_truncate_codes(codes, expected);
let codes_tensor =
Tensor::<B, 2, Int>::from_data(TensorData::new(codes, [1, expected]), &device);
let hints = self.condition.codes_to_hints(codes_tensor);
let [_, hint_frames, _] = hints.dims();
let src_latents = if hint_frames >= latent_frames {
hints.narrow(1, 0, latent_frames)
} else {
let padding = self.silence.slice(latent_frames - hint_frames);
Tensor::cat(vec![hints, padding], 1)
};
if let Some(trace) = trace.as_deref_mut() {
trace.hints_latent_rms = tensor_rms(&src_latents)?;
progress("vae", 0.0, "decoding hint latents (trace)");
let hints_audio = self.vae.decode(src_latents.clone());
let interleaved = interleave_channel_major(hints_audio)?;
trace.hints_audio = Some(interleaved);
progress("vae", 0.3, "hint latents decoded (trace)");
}
let src_latents = if std::env::var_os("MAOLAN_ACESTEP_NOCODES").is_some() {
self.silence.slice(latent_frames)
} else {
src_latents
};
let chunk_mask = Tensor::ones([1, latent_frames, acoustic_dim], &device);
let context = Tensor::cat(vec![src_latents, chunk_mask], 2);
let noise = seeded_latent_noise(seed, latent_frames, acoustic_dim, &device);
let timesteps: &[f32] = if self.dit_config.is_turbo {
&TURBO_TIMESTEPS
} else {
&SFT_TIMESTEPS
};
let latents = if let Some(trace) = trace.as_deref_mut() {
let kv = self.dit.prepare_cross_kv(enc);
let total = timesteps.len();
let mut xt = noise;
for (index, &t_cur) in timesteps.iter().enumerate() {
let v = self
.dit
.forward_with_kv(xt.clone(), t_cur, context.clone(), &kv);
trace.dit_steps.push(DitStepStat {
t: t_cur,
xt_rms: tensor_rms(&xt)?,
v_rms: tensor_rms(&v)?,
});
let dt = if index + 1 == total {
t_cur
} else {
t_cur - timesteps[index + 1]
};
xt = xt - v * dt;
progress("dit", (index + 1) as f32 / total as f32, "diffusing");
}
xt
} else {
let mut dit_progress = |done: usize, total: usize| {
progress("dit", done as f32 / total as f32, "diffusing");
};
self.dit
.sample_turbo(noise, context, enc, timesteps, Some(&mut dit_progress))
};
if let Some(trace) = trace {
trace.final_latent_rms = tensor_rms(&latents)?;
}
progress("vae", 0.0, "decoding audio");
let audio = self.vae.decode(latents);
progress("vae", 1.0, "audio decoded");
let [_, channels, frames] = audio.dims();
let meta = GenerateAudioMeta {
channels,
frames,
sample_rate_hz: self.vae_config.sampling_rate as u32,
steps: timesteps.len(),
prompt_tokens,
};
Ok((audio, meta))
}
}
fn latent_frames_for_duration(duration_ms: usize) -> usize {
((duration_ms * 25 + 500) / 1000).max(5)
}
fn tensor_rms<B: Backend>(tensor: &Tensor<B, 3>) -> Result<f32> {
let values: Vec<f32> = tensor
.clone()
.into_data()
.convert::<f32>()
.to_vec()
.map_err(|e| anyhow::anyhow!("failed to read tensor: {e}"))?;
Ok((values.iter().map(|v| v * v).sum::<f32>() / values.len() as f32).sqrt())
}
fn interleave_channel_major<B: Backend>(audio: Tensor<B, 3>) -> Result<(Vec<f32>, usize, usize)> {
let [_, channels, frames] = audio.dims();
let channel_major: Vec<f32> = audio
.into_data()
.convert::<f32>()
.to_vec()
.map_err(|e| anyhow::anyhow!("failed to read audio tensor: {e}"))?;
let mut interleaved = vec![0.0_f32; channel_major.len()];
for (channel, samples) in channel_major.chunks_exact(frames).enumerate() {
for (frame, sample) in samples.iter().enumerate() {
interleaved[frame * channels + channel] = *sample;
}
}
Ok((interleaved, channels, frames))
}
const MAX_TEXT_TOKENS: usize = 256;
const MAX_LYRIC_TOKENS: usize = 2048;
pub const DIT_INSTRUCTION: &str = "Generate audio semantic tokens based on the given conditions:";
pub const INSTRUMENTAL_LYRIC_PROMPT: &str =
"# Languages\nunknown\n\n# Lyric\n[Instrumental]<|endoftext|>";
pub fn build_metas_block(
bpm: Option<f32>,
key_scale: Option<&str>,
time_signature: Option<&str>,
duration_s: usize,
) -> String {
let bpm = match bpm {
Some(bpm) if bpm.fract() == 0.0 => format!("{}", bpm as i64),
Some(bpm) => format!("{bpm}"),
None => "N/A".to_string(),
};
let time_signature = time_signature.unwrap_or("N/A");
let key_scale = key_scale.unwrap_or("N/A");
format!(
"- bpm: {bpm}\n- timesignature: {time_signature}\n- keyscale: {key_scale}\n- duration: {duration_s} seconds\n"
)
}
pub fn build_dit_text_prompt(caption: &str, metas_block: &str) -> String {
format!(
"# Instruction\n{DIT_INSTRUCTION}\n\n# Caption\n{caption}\n\n# Metas\n{metas_block}<|endoftext|>\n"
)
}
fn pad_or_truncate_codes(mut codes: Vec<u32>, expected: usize) -> Vec<u32> {
if codes.len() < expected {
let fill = codes.last().copied().unwrap_or(0);
codes.resize(expected, fill);
} else {
codes.truncate(expected);
}
codes
}
pub fn tile_frames<B: Backend>(tensor: Tensor<B, 3>, frames: usize) -> Tensor<B, 3> {
let len = tensor.dims()[1];
assert!(len > 0, "cannot tile an empty latent");
if len == frames {
return tensor;
}
let repeats = frames.div_ceil(len);
let tiled = if repeats > 1 {
tensor.repeat_dim(1, repeats)
} else {
tensor
};
tiled.narrow(1, 0, frames)
}
pub fn seeded_latent_noise<B: Backend>(
seed: u64,
frames: usize,
channels: usize,
device: &B::Device,
) -> Tensor<B, 3> {
let data = generate_gaussian_data(seed, frames * channels);
Tensor::<B, 3>::from_data(TensorData::new(data, [1, frames, channels]), device)
}
fn generate_gaussian_data(seed: u64, len: usize) -> Vec<f32> {
let mut out = Vec::with_capacity(len);
let mut state = seed;
while out.len() < len {
let u1 = uniform01_open(&mut state);
let u2 = uniform01_open(&mut state);
let radius = (-2.0_f64 * u1.ln()).sqrt();
let theta = 2.0_f64 * std::f64::consts::PI * u2;
out.push((radius * theta.cos()) as f32);
if out.len() < len {
out.push((radius * theta.sin()) as f32);
}
}
out
}
fn uniform01_open(state: &mut u64) -> f64 {
let value = splitmix64_next(state);
let mantissa = (value >> 11) as f64;
((mantissa + 0.5) / ((1_u64 << 53) as f64)).clamp(f64::MIN_POSITIVE, 1.0 - f64::EPSILON)
}
fn splitmix64_next(state: &mut u64) -> u64 {
*state = state.wrapping_add(0x9E3779B97F4A7C15);
let mut z = *state;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58476D1CE4E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D049BB133111EB);
z ^ (z >> 31)
}
#[cfg(test)]
mod tests {
use super::*;
use burn::backend::NdArray;
type TestBackend = NdArray<f32>;
#[test]
fn metas_block_matches_official_format() {
assert_eq!(
build_metas_block(Some(120.0), Some("A minor"), Some("4/4"), 10),
"- bpm: 120\n- timesignature: 4/4\n- keyscale: A minor\n- duration: 10 seconds\n"
);
assert_eq!(
build_metas_block(None, None, None, 30),
"- bpm: N/A\n- timesignature: N/A\n- keyscale: N/A\n- duration: 30 seconds\n"
);
assert_eq!(
build_metas_block(Some(128.5), None, Some("6/8"), 5),
"- bpm: 128.5\n- timesignature: 6/8\n- keyscale: N/A\n- duration: 5 seconds\n"
);
}
#[test]
fn dit_text_prompt_matches_official_sft_format() {
let metas = build_metas_block(Some(120.0), Some("A minor"), Some("4/4"), 10);
assert_eq!(
build_dit_text_prompt("dark techno", &metas),
"# Instruction\nGenerate audio semantic tokens based on the given conditions:\n\n# Caption\ndark techno\n\n# Metas\n- bpm: 120\n- timesignature: 4/4\n- keyscale: A minor\n- duration: 10 seconds\n<|endoftext|>\n"
);
}
#[test]
fn instrumental_lyric_prompt_matches_official_format() {
assert_eq!(
INSTRUMENTAL_LYRIC_PROMPT,
"# Languages\nunknown\n\n# Lyric\n[Instrumental]<|endoftext|>"
);
}
#[test]
fn latent_frame_count_math() {
assert_eq!(latent_frames_for_duration(0), 5, "minimum of 5 frames");
assert_eq!(latent_frames_for_duration(100), 5);
assert_eq!(latent_frames_for_duration(1000), 25);
assert_eq!(latent_frames_for_duration(30_000), 750);
assert_eq!(latent_frames_for_duration(1020), 26);
}
#[test]
fn pad_or_truncate_codes_behaviour() {
assert_eq!(pad_or_truncate_codes(vec![], 3), vec![0, 0, 0]);
assert_eq!(pad_or_truncate_codes(vec![7, 9], 4), vec![7, 9, 9, 9]);
assert_eq!(pad_or_truncate_codes(vec![1, 2, 3, 4], 2), vec![1, 2]);
assert_eq!(pad_or_truncate_codes(vec![5], 1), vec![5]);
}
#[test]
fn tile_frames_tiles_and_crops() {
let device = Default::default();
let tensor = Tensor::<TestBackend, 3>::from_data(
[[
[1.0, 10.0, 100.0, 1000.0],
[2.0, 20.0, 200.0, 2000.0],
[3.0, 30.0, 300.0, 3000.0],
]],
&device,
);
let tiled = tile_frames(tensor.clone(), 7);
assert_eq!(tiled.dims(), [1, 7, 4]);
let values = tiled.to_data().to_vec::<f32>().expect("tiled values");
for frame in 0..7 {
let source = frame % 3;
assert_eq!(values[frame * 4], (source + 1) as f32);
}
let cropped = tile_frames(tensor.clone(), 2);
assert_eq!(cropped.dims(), [1, 2, 4]);
let values = cropped.to_data().to_vec::<f32>().expect("cropped values");
assert_eq!(values[0], 1.0);
assert_eq!(values[4], 2.0);
assert_eq!(tile_frames(tensor, 3).dims(), [1, 3, 4]);
}
#[test]
fn silence_latent_slice_and_timbre_reference_tile() {
let device = Default::default();
let tensor = Tensor::<TestBackend, 3>::from_data(
[[
[1.0, 2.0, 3.0, 4.0],
[5.0, 6.0, 7.0, 8.0],
[9.0, 10.0, 11.0, 12.0],
]],
&device,
);
let silence = SilenceLatent::<TestBackend> {
silence_latent: Param::from_tensor(tensor),
};
let sliced = silence.slice(5);
assert_eq!(sliced.dims(), [1, 5, 4]);
let values = sliced.to_data().to_vec::<f32>().expect("slice values");
assert_eq!(&values[0..4], &[1.0, 2.0, 3.0, 4.0]);
assert_eq!(&values[12..16], &[1.0, 2.0, 3.0, 4.0]);
assert_eq!(&values[16..20], &[5.0, 6.0, 7.0, 8.0]);
let reference = silence.timbre_reference();
assert_eq!(reference.dims(), [1, TIMBRE_REFERENCE_FRAMES, 4]);
let values = reference
.to_data()
.to_vec::<f32>()
.expect("reference values");
let last = &values[(TIMBRE_REFERENCE_FRAMES - 1) * 4..];
assert_eq!(last, &[9.0, 10.0, 11.0, 12.0]);
}
#[test]
fn seeded_noise_is_deterministic_and_normal_shaped() {
let device = Default::default();
let first = seeded_latent_noise::<TestBackend>(42, 6, 4, &device);
let second = seeded_latent_noise::<TestBackend>(42, 6, 4, &device);
assert_eq!(first.dims(), [1, 6, 4]);
let a = first.to_data().to_vec::<f32>().expect("first noise");
let b = second.to_data().to_vec::<f32>().expect("second noise");
assert_eq!(a, b, "same seed must reproduce the same noise");
assert!(a.iter().all(|v| v.is_finite()));
assert!(
a.iter().any(|v| v.abs() > 1e-6),
"noise must not be all zeros"
);
let other = seeded_latent_noise::<TestBackend>(7, 6, 4, &device)
.to_data()
.to_vec::<f32>()
.expect("other noise");
assert_ne!(a, other, "different seeds must diverge");
}
}