use anyhow::Result;
use crate::kv_cache::InferenceState;
use crate::model::Model;
use crate::model::audio_decoder::{
AudioDecoderWeights, AudioGpu, DepthformerState, DetokenizerState, DetokenizerWeights,
detokenize_to_spectrum, embed_audio_token, istft_to_pcm, sample_audio_frame,
};
use crate::sampler::{Sampler, SamplerConfig};
use crate::time::{Duration, Instant};
use crate::tokenizer::BpeTokenizer;
pub struct AudioGenerateConfig {
pub max_tokens: usize,
pub sampler: SamplerConfig,
pub audio_temperature: f32,
pub audio_top_k: usize,
pub mode: AudioMode,
pub gpu_depthformer: bool,
}
#[derive(Clone, Copy)]
pub enum AudioMode {
Sequential,
Interleaved,
}
pub struct AudioGenerateResult {
pub text_tokens: usize,
pub audio_frames: usize,
pub audio_samples: usize,
pub elapsed_secs: f64,
pub depthformer_secs: f64,
pub detokenizer_secs: f64,
}
const TOKEN_AUDIO_START: u32 = 128;
const TOKEN_TEXT_END: u32 = 130;
const AUDIO_END_CODE: i32 = 2048;
#[derive(PartialEq)]
enum Modality {
Text,
Audio,
}
enum FrameOutcome {
Codes { audio_embedding: Vec<f32> },
End,
}
struct AudioOutputDecoder<'a> {
weights: &'a AudioDecoderWeights,
detok_weights: &'a DetokenizerWeights,
gpu: Option<&'a dyn AudioGpu>,
df_state: DepthformerState,
detok_state: DetokenizerState,
all_spectrum: Vec<f32>,
audio_frames: usize,
time_depthformer: Duration,
time_detokenizer: Duration,
audio_temperature: f32,
audio_top_k: usize,
use_gpu_df: bool,
}
impl<'a> AudioOutputDecoder<'a> {
fn new(
weights: &'a AudioDecoderWeights,
detok_weights: &'a DetokenizerWeights,
gpu: Option<&'a dyn AudioGpu>,
audio_temperature: f32,
audio_top_k: usize,
gpu_depthformer: bool,
) -> Self {
if let Some(g) = gpu {
g.reset_detokenizer();
g.reset_depthformer();
}
let df_state = DepthformerState::new(&weights.depthformer_config);
let detok_state = DetokenizerState::new(&detok_weights.config);
Self {
weights,
detok_weights,
gpu,
df_state,
detok_state,
all_spectrum: Vec::new(),
audio_frames: 0,
time_depthformer: Duration::ZERO,
time_detokenizer: Duration::ZERO,
audio_temperature,
audio_top_k,
use_gpu_df: gpu_depthformer && gpu.is_some(),
}
}
fn decode_frame(&mut self, embed: &[f32]) -> FrameOutcome {
let t0 = Instant::now();
let codes = if self.use_gpu_df {
let g = self
.gpu
.expect("use_gpu_df implies gpu is Some (set at construction)");
g.sample_audio_frame(embed, self.audio_temperature, self.audio_top_k)
} else {
sample_audio_frame(
self.weights,
&mut self.df_state,
embed,
self.audio_temperature,
self.audio_top_k,
)
};
self.time_depthformer += t0.elapsed();
if codes[0] == AUDIO_END_CODE {
return FrameOutcome::End;
}
let t1 = Instant::now();
let spectrum = if let Some(g) = self.gpu {
g.detokenize_to_spectrum(self.detok_weights, &codes)
} else {
detokenize_to_spectrum(
self.detok_weights,
self.weights,
&mut self.detok_state,
&codes,
)
};
self.time_detokenizer += t1.elapsed();
self.all_spectrum.extend_from_slice(&spectrum);
self.audio_frames += 1;
let audio_embedding = embed_audio_token(self.weights, &codes);
FrameOutcome::Codes { audio_embedding }
}
fn finish(&mut self, mut sink: impl FnMut(&[f32], u32)) -> usize {
if self.all_spectrum.is_empty() {
return 0;
}
let pcm = istft_to_pcm(
&self.all_spectrum,
self.detok_weights.config.n_fft,
self.detok_weights.config.hop_length,
);
if pcm.is_empty() {
return 0;
}
let n = pcm.len();
sink(&pcm, self.detok_weights.config.sample_rate as u32);
n
}
}
#[allow(unused_assignments, clippy::too_many_arguments)]
pub fn generate_audio(
model: &dyn Model,
decoder_weights: &AudioDecoderWeights,
detok_weights: &DetokenizerWeights,
tokenizer: &BpeTokenizer,
prompt_tokens: &[u32],
config: &AudioGenerateConfig,
gpu: Option<&dyn AudioGpu>,
mut text_callback: impl FnMut(&str),
mut audio_callback: impl FnMut(&[f32], u32),
) -> Result<AudioGenerateResult> {
anyhow::ensure!(!prompt_tokens.is_empty(), "prompt_tokens must not be empty");
let model_config = model.config();
let mut state = InferenceState::from_config(model_config)?;
let mut sampler = Sampler::new(config.sampler.clone());
let mut decoder = AudioOutputDecoder::new(
decoder_weights,
detok_weights,
gpu,
config.audio_temperature,
config.audio_top_k,
config.gpu_depthformer,
);
let start = Instant::now();
let mut logits = model.forward_prefill(prompt_tokens, 0, &mut state);
let mut modality = Modality::Text;
let mut generated = 0usize;
let mut text_tokens = 0usize;
let mut pos = prompt_tokens.len();
let mut modality_budget = match config.mode {
AudioMode::Interleaved => 6, AudioMode::Sequential => usize::MAX,
};
let mut text_done = false;
let mut next_token = sampler.sample(&mut logits);
let mut trailing_audio_segments: usize = 0;
const MAX_TRAILING_AUDIO_SEGMENTS: usize = 3;
'outer: loop {
if generated >= config.max_tokens || pos >= model_config.max_seq_len {
break;
}
if modality == Modality::Text {
if tokenizer.eos_token() == Some(next_token) {
break;
}
if next_token == TOKEN_AUDIO_START {
modality = Modality::Audio;
modality_budget = match config.mode {
AudioMode::Interleaved => 12,
AudioMode::Sequential => usize::MAX,
};
continue;
}
if next_token == TOKEN_TEXT_END {
text_done = true;
}
if next_token != TOKEN_TEXT_END {
let piece = tokenizer.decode(&[next_token]);
text_callback(&piece);
text_tokens += 1;
}
generated += 1;
modality_budget = modality_budget.saturating_sub(1);
if generated >= config.max_tokens {
break;
}
if matches!(config.mode, AudioMode::Interleaved) && (modality_budget == 0 || text_done)
{
if text_done {
trailing_audio_segments += 1;
if trailing_audio_segments > MAX_TRAILING_AUDIO_SEGMENTS {
break;
}
}
let mut emb = model.forward_embedding(&[next_token], pos, &mut state);
pos += 1;
modality = Modality::Audio;
modality_budget = 12;
loop {
let outcome = decoder.decode_frame(&emb);
let audio_emb = match outcome {
FrameOutcome::End => {
text_done = true;
break;
}
FrameOutcome::Codes { audio_embedding } => audio_embedding,
};
modality_budget = modality_budget.saturating_sub(1);
if generated >= config.max_tokens || pos >= model_config.max_seq_len {
break;
}
if modality_budget == 0 && !text_done {
logits = model.forward_from_embedding(&audio_emb, pos, &mut state);
next_token = sampler.sample(&mut logits);
pos += 1;
break;
}
emb = model.forward_hidden_from_embedding(&audio_emb, pos, &mut state);
pos += 1;
generated += 1;
}
modality = Modality::Text;
modality_budget = 6;
continue;
}
logits = model.forward(&[next_token], pos, &mut state);
next_token = sampler.sample(&mut logits);
pos += 1;
} else {
let mut emb = model.forward_embedding(&[next_token], pos, &mut state);
pos += 1;
generated += 1;
loop {
let outcome = decoder.decode_frame(&emb);
let audio_emb = match outcome {
FrameOutcome::End => match config.mode {
AudioMode::Sequential => {
break 'outer;
}
AudioMode::Interleaved => {
modality = Modality::Text;
text_done = true;
modality_budget = 6;
logits = model.forward(&[TOKEN_TEXT_END], pos, &mut state);
next_token = sampler.sample(&mut logits);
pos += 1;
break;
}
},
FrameOutcome::Codes { audio_embedding } => audio_embedding,
};
modality_budget = modality_budget.saturating_sub(1);
emb = model.forward_hidden_from_embedding(&audio_emb, pos, &mut state);
pos += 1;
generated += 1;
if generated >= config.max_tokens || pos >= model_config.max_seq_len {
break;
}
if matches!(config.mode, AudioMode::Interleaved)
&& modality_budget == 0
&& !text_done
{
modality = Modality::Text;
modality_budget = 6;
logits = model.forward_from_embedding(&audio_emb, pos, &mut state);
next_token = sampler.sample(&mut logits);
pos += 1;
break;
}
}
}
}
let audio_samples = decoder.finish(&mut audio_callback);
Ok(AudioGenerateResult {
text_tokens,
audio_frames: decoder.audio_frames,
audio_samples,
elapsed_secs: start.elapsed().as_secs_f64(),
depthformer_secs: decoder.time_depthformer.as_secs_f64(),
detokenizer_secs: decoder.time_detokenizer.as_secs_f64(),
})
}