diff --git a/Cargo.lock b/Cargo.lock index 64a2b97..09bcd9f 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -16,7 +16,7 @@ checksum = "b169f7a6d4742236a0a00c541b845991d0ac43e546831af1249753ab4c3aa3a0" dependencies = [ "cfg-if", "cipher", - "cpufeatures", + "cpufeatures 0.2.17", ] [[package]] @@ -42,6 +42,7 @@ dependencies = [ "minijinja", "modelscope", "num", + "rand 0.10.1", "rayon", "realfft", "reqwest", @@ -662,6 +663,17 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724" +[[package]] +name = "chacha20" +version = "0.10.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6f8d983286843e49675a4b7a2d174efe136dc93a18d69130dd18198a6c167601" +dependencies = [ + "cfg-if", + "cpufeatures 0.3.0", + "rand_core 0.10.1", +] + [[package]] name = "cipher" version = "0.4.4" @@ -871,6 +883,15 @@ dependencies = [ "libc", ] +[[package]] +name = "cpufeatures" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b2a41393f66f16b0823bb79094d54ac5fbd34ab292ddafb9a0456ac9f87d201" +dependencies = [ + "libc", +] + [[package]] name = "crc" version = "3.3.0" @@ -1818,6 +1839,7 @@ dependencies = [ "cfg-if", "libc", "r-efi 6.0.0", + "rand_core 0.10.1", "wasip2", "wasip3", ] @@ -3400,6 +3422,17 @@ dependencies = [ "rand_core 0.9.5", ] +[[package]] +name = "rand" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d2e8e8bcc7961af1fdac401278c6a831614941f6164ee3bf4ce61b7edb162207" +dependencies = [ + "chacha20", + "getrandom 0.4.2", + "rand_core 0.10.1", +] + [[package]] name = "rand_chacha" version = "0.3.1" @@ -3438,6 +3471,12 @@ dependencies = [ "getrandom 0.3.4", ] +[[package]] +name = "rand_core" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "63b8176103e19a2643978565ca18b50549f6101881c443590420e4dc998a3c69" + [[package]] name = "rand_distr" version = "0.4.3" @@ -4093,7 +4132,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e3bf829a2d51ab4a5ddf1352d8470c140cadc8301b2ae1789db023f01cedd6ba" dependencies = [ "cfg-if", - "cpufeatures", + "cpufeatures 0.2.17", "digest", ] @@ -4104,7 +4143,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283" dependencies = [ "cfg-if", - "cpufeatures", + "cpufeatures 0.2.17", "digest", ] diff --git a/Cargo.toml b/Cargo.toml index e122bce..a7222c6 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -45,6 +45,7 @@ sentencepiece = "0.13.1" ahash = "0.8.12" derive_builder = "0.20.2" kaldi-native-fbank = "0.1.0" +rand = "0.10.1" [patch.crates-io] esaxx-rs = { path = "vendor/esaxx-rs" } diff --git a/assets/audio/jiangjiang.wav b/assets/audio/jiangjiang.wav new file mode 100644 index 0000000..cc827d0 Binary files /dev/null and b/assets/audio/jiangjiang.wav differ diff --git a/assets/audio/voice_01.wav b/assets/audio/voice_01.wav deleted file mode 100644 index 224a4e9..0000000 Binary files a/assets/audio/voice_01.wav and /dev/null differ diff --git a/assets/audio/voice_05.wav b/assets/audio/voice_05.wav deleted file mode 100644 index 2dcc49d..0000000 Binary files a/assets/audio/voice_05.wav and /dev/null differ diff --git a/assets/audio/zh.mp3 b/assets/audio/zh.mp3 deleted file mode 100644 index 1ae2c89..0000000 Binary files a/assets/audio/zh.mp3 and /dev/null differ diff --git a/src/models/bigvgan/mod.rs b/src/models/bigvgan/mod.rs index 20ffa41..8ee1767 100644 --- a/src/models/bigvgan/mod.rs +++ b/src/models/bigvgan/mod.rs @@ -186,6 +186,8 @@ impl AMPBlock1 { 1, 1, true, + None, + None, )?; convs1.push(layer); } @@ -203,6 +205,8 @@ impl AMPBlock1 { 1, 1, true, + None, + None, )?; convs2.push(layer); } @@ -261,6 +265,8 @@ impl BigVGAN { 1, 1, true, + None, + None, )?; let vb_ups = vb.pp("ups"); @@ -296,7 +302,7 @@ impl BigVGAN { } } let activation_post = TorchActivation1d::new(vb.pp("activation_post"), 2, 2, 12, 12, ch)?; - let conv_post = WNConv1d::new(vb.pp("conv_post"), ch, 1, 7, 1, 3, 1, 1, false)?; + let conv_post = WNConv1d::new(vb.pp("conv_post"), ch, 1, 7, 1, 3, 1, 1, false, None, None)?; Ok(Self { num_kernels, diff --git a/src/models/common/generate.rs b/src/models/common/generate.rs index cd3502d..f9383d3 100644 --- a/src/models/common/generate.rs +++ b/src/models/common/generate.rs @@ -1,12 +1,12 @@ use anyhow::Result; use candle_core::{DType, Device, Tensor}; -use candle_transformers::generation::{LogitsProcessor, Sampling}; +use candle_transformers::generation::{LogitsProcessor}; use rocket::async_stream::stream; use rocket::futures::Stream; use std::time::Instant; use crate::{ - models::common::{InferenceModel, MultiModalData}, + models::common::{InferenceModel, MultiModalData, sample::{use_repeat_penalty, get_logit_processor}}, params::chat::{ChatCompletionChunkResponse, ChatCompletionResponse}, tokenizer::TokenizerModel, utils::response_utils::{ @@ -14,38 +14,6 @@ use crate::{ build_completion_chunk_response, build_completion_response_with_time, }, }; -pub fn get_logit_processor( - temperature: Option, - top_p: Option, - top_k: Option, - seed: u64, -) -> LogitsProcessor { - let temperature = temperature.and_then(|v| if v < 1e-7 { None } else { Some(v) }); - match top_k { - None => LogitsProcessor::new( - seed, - temperature.map(|temp| temp as f64), - top_p.map(|tp| tp as f64), - ), - Some(k) => { - let sampling = match temperature { - None => Sampling::ArgMax, - Some(temperature) => match top_p { - None => Sampling::TopK { - k, - temperature: temperature as f64, - }, - Some(p) => Sampling::TopKThenTopP { - k, - p: p as f64, - temperature: temperature as f64, - }, - }, - }; - LogitsProcessor::from_sampling(seed, sampling) - } - } -} pub struct GenerationContext { pub logit_processor: LogitsProcessor, @@ -103,16 +71,12 @@ fn sample_and_push( ) -> Result { let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; // 重复惩罚 - let logits = if ctx.repeat_penalty == 1. || ctx.repeat_last_n == 0 { - logits - } else { - let start_at = generated.len().saturating_sub(ctx.repeat_last_n); - candle_transformers::utils::apply_repeat_penalty( - &logits, - ctx.repeat_penalty, - &generated[start_at..], - )? - }; + let logits = use_repeat_penalty( + ctx.repeat_penalty, + Some(ctx.repeat_last_n), + &logits, + generated, + )?; let token = ctx.logit_processor.sample(&logits)?; generated.push(token); Ok(token) diff --git a/src/models/common/mod.rs b/src/models/common/mod.rs index 5286926..c2f6f5f 100644 --- a/src/models/common/mod.rs +++ b/src/models/common/mod.rs @@ -6,6 +6,7 @@ pub mod gguf; pub mod model_mapping; pub mod modules; pub mod reranker; +pub mod sample; /// 多模态模型的特征数据 /// 每个模型数据不一样 diff --git a/src/models/common/modules.rs b/src/models/common/modules.rs index 1f335d0..b0f877e 100644 --- a/src/models/common/modules.rs +++ b/src/models/common/modules.rs @@ -2,14 +2,14 @@ use anyhow::{Result, anyhow}; use candle_core::{D, DType, IndexOp, Tensor}; use candle_nn::{ Activation, BatchNorm, BatchNormConfig, Conv1d, Conv1dConfig, Conv2d, Conv2dConfig, - ConvTranspose1d, ConvTranspose1dConfig, Embedding, LayerNorm, LayerNormConfig, Linear, Module, - ModuleT, RmsNorm, VarBuilder, batch_norm, conv1d, conv1d_no_bias, conv2d, conv2d_no_bias, - embedding, layer_norm, linear_b, linear_no_bias, ops::sigmoid, rms_norm, + ConvTranspose1d, ConvTranspose1dConfig, LayerNorm, LayerNormConfig, Linear, Module, ModuleT, + RmsNorm, VarBuilder, batch_norm, conv1d, conv1d_no_bias, conv2d, conv2d_no_bias, layer_norm, + linear_b, ops::sigmoid, rms_norm, }; use crate::{ - position_embed::rope::{RoPE, apply_rotary_pos_emb, apply_rotary_pos_emb_roformer}, - utils::tensor_utils::{pad_replicate_last_dim, prepare_causal_attention_mask, repeat_kv}, + position_embed::rope::{apply_rotary_pos_emb, apply_rotary_pos_emb_roformer}, + utils::tensor_utils::{pad_replicate_last_dim, repeat_kv}, }; #[derive(Debug)] @@ -371,7 +371,7 @@ impl QKVCatAttention { && let Some(sin) = sin { if use_roformer { - apply_rotary_pos_emb_roformer(&query_states, &key_states, cos, sin, tof32)? + apply_rotary_pos_emb_roformer(&query_states, &key_states, cos, sin)? } else { apply_rotary_pos_emb(&query_states, &key_states, cos, sin, tof32)? } @@ -412,7 +412,7 @@ impl QKVCatAttention { let key_states = qkv.i(1)?.contiguous()?; let value_states = qkv.i(2)?.contiguous()?; let (query_states, key_states) = if use_roformer { - apply_rotary_pos_emb_roformer(&query_states, &key_states, cos, sin, tof32)? + apply_rotary_pos_emb_roformer(&query_states, &key_states, cos, sin)? } else { apply_rotary_pos_emb(&query_states, &key_states, cos, sin, tof32)? }; @@ -765,11 +765,11 @@ pub fn eager_attention_forward( // input q shape:(b, num_head, seq_len, dim) // input k/v shape:(b, num_kv_head, seq_len, dim) let key_states = match num_key_value_groups { - Some(g) => repeat_kv(key_states.clone(), g)?.contiguous()?, + Some(g) => repeat_kv(key_states.clone(), g)?, None => key_states.clone(), }; let value_states = match num_key_value_groups { - Some(g) => repeat_kv(value_states.clone(), g)?.contiguous()?, + Some(g) => repeat_kv(value_states.clone(), g)?, None => value_states.clone(), }; let query_states = query_states.contiguous()?; @@ -778,13 +778,14 @@ pub fn eager_attention_forward( let attn_output = { #[cfg(not(feature = "flash-attn"))] { - let attn_weights = query_states.matmul(&key_states.transpose(D::Minus2, D::Minus1)?)?; + let attn_weights = + query_states.matmul(&key_states.transpose(D::Minus2, D::Minus1)?.contiguous()?)?; let attn_weights = (attn_weights * scaling)?; let attn_weights = match attention_mask { None => attn_weights, Some(mask) => attn_weights.broadcast_add(&mask.to_dtype(attn_weights.dtype())?)?, }; - let attn_weights = candle_nn::ops::softmax_last_dim(&attn_weights)?; + let attn_weights = candle_nn::ops::softmax_last_dim(&attn_weights)?.contiguous()?; attn_weights.matmul(&value_states)? } #[cfg(feature = "flash-attn")] @@ -987,163 +988,6 @@ pub fn deform_conv2d_kernel( Ok(out) } -pub struct LlamaModel { - pub embed_tokens: Embedding, - layers: Vec, - norm: RmsNorm, - rotary_emb: RoPE, -} - -impl LlamaModel { - pub fn new( - vb: VarBuilder, - vocab_size: usize, - hidden_size: usize, - num_hidden_layers: usize, - num_attention_heads: usize, - num_key_value_heads: Option, - head_dim: Option, - attn_bias: bool, - attn_pp_name: &str, - o_proj_pp_name: Option<&str>, - intermediate_size: usize, - hidden_act: Activation, - mlp_bias: bool, - mlp_pp_name: &str, - norm_eps: f64, - input_norm_pp_name: &str, - post_norm_pp_name: &str, - rope_theta_base: f32, - ) -> Result { - let embed_tokens = embedding(vocab_size, hidden_size, vb.pp("embed_tokens"))?; - let mut layers = vec![]; - let vb_layers = vb.pp("layers"); - for i in 0..num_hidden_layers { - let layers_i = NaiveAttnGateUpDownMLPBlock::new( - vb_layers.pp(i), - hidden_size, - num_attention_heads, - num_key_value_heads, - head_dim, - attn_bias, - attn_pp_name, - o_proj_pp_name, - intermediate_size, - hidden_act, - mlp_bias, - mlp_pp_name, - norm_eps, - input_norm_pp_name, - post_norm_pp_name, - )?; - layers.push(layers_i); - } - let norm = rms_norm(hidden_size, norm_eps, vb.pp("norm"))?; - let head_dim = head_dim.unwrap_or(hidden_size / num_attention_heads); - let rotary_emb = RoPE::new(head_dim, rope_theta_base, vb.device())?; - Ok(Self { - embed_tokens, - layers, - norm, - rotary_emb, - }) - } - - pub fn forward(&mut self, inputs_embeds: &Tensor, seqlen_offset: usize) -> Result { - let (b_size, seq_len, _) = inputs_embeds.dims3()?; - - let (cos, sin) = self - .rotary_emb - .forward(seqlen_offset, seq_len, inputs_embeds.device())?; - let mut xs = inputs_embeds.clone(); - let attention_mask: Option = { - if seq_len <= 1 { - None - } else { - Some(prepare_causal_attention_mask( - b_size, - seq_len, - 0, - xs.device(), - )?) - } - }; - for layer in self.layers.iter_mut() { - xs = layer.forward(&xs, &cos, &sin, attention_mask.as_ref())?; - } - let xs = xs.apply(&self.norm)?; - Ok(xs) - } - - pub fn clear_kv_cache(&mut self) { - for layer in self.layers.iter_mut() { - layer.clear_kv_cache() - } - } -} - -pub struct LlamaForCausalLM { - pub model: LlamaModel, - lm_head: Linear, -} - -impl LlamaForCausalLM { - pub fn new( - vb: VarBuilder, - vocab_size: usize, - hidden_size: usize, - num_hidden_layers: usize, - num_attention_heads: usize, - num_key_value_heads: Option, - head_dim: Option, - attn_bias: bool, - attn_pp_name: &str, - o_proj_pp_name: Option<&str>, - intermediate_size: usize, - hidden_act: Activation, - mlp_bias: bool, - mlp_pp_name: &str, - norm_eps: f64, - input_norm_pp_name: &str, - post_norm_pp_name: &str, - rope_theta_base: f32, - ) -> Result { - let model = LlamaModel::new( - vb.pp("model"), - vocab_size, - hidden_size, - num_hidden_layers, - num_attention_heads, - num_key_value_heads, - head_dim, - attn_bias, - attn_pp_name, - o_proj_pp_name, - intermediate_size, - hidden_act, - mlp_bias, - mlp_pp_name, - norm_eps, - input_norm_pp_name, - post_norm_pp_name, - rope_theta_base, - )?; - let lm_head = linear_no_bias(hidden_size, vocab_size, vb.pp("lm_head"))?; - Ok(Self { model, lm_head }) - } - - pub fn forward(&mut self, inputs_embeds: &Tensor, seqlen_offset: usize) -> Result { - let outputs = self.model.forward(inputs_embeds, seqlen_offset)?; - let seq_len = outputs.dim(1)?; - let hidden_state = outputs.narrow(1, seq_len - 1, 1)?; - let logits = self.lm_head.forward(&hidden_state)?; - Ok(logits) - } - pub fn clear_kv_cache(&mut self) { - self.model.clear_kv_cache(); - } -} - pub struct GLU { dim: usize, } @@ -1190,10 +1034,14 @@ impl WNConv1d { groups: usize, stride: usize, bias: bool, + weight_g_pp_name: Option<&str>, + weight_v_pp_name: Option<&str>, ) -> Result { let in_c = in_c / groups; - let weight_g = vb.get((out_c, 1, 1), "weight_g")?; - let weight_v = vb.get((out_c, in_c, kernel_size), "weight_v")?; + let weight_g_pp_name = weight_g_pp_name.unwrap_or("weight_g"); + let weight_v_pp_name = weight_v_pp_name.unwrap_or("weight_v"); + let weight_g = vb.get((out_c, 1, 1), weight_g_pp_name)?; + let weight_v = vb.get((out_c, in_c, kernel_size), weight_v_pp_name)?; // let bias = vb.get(out_c, "bias").ok(); let bias = if bias { vb.get(out_c, "bias").ok() diff --git a/src/models/common/sample.rs b/src/models/common/sample.rs new file mode 100644 index 0000000..7f117ac --- /dev/null +++ b/src/models/common/sample.rs @@ -0,0 +1,140 @@ +use anyhow::{Result, anyhow}; +use candle_core::{IndexOp, Tensor}; +use candle_nn::ops::softmax; +use candle_transformers::generation::{LogitsProcessor, Sampling}; +use rand::{SeedableRng, distr::Distribution}; + +pub fn get_logit_processor( + temperature: Option, + top_p: Option, + top_k: Option, + seed: u64, +) -> LogitsProcessor { + let temperature = temperature.and_then(|v| if v < 1e-7 { None } else { Some(v) }); + match top_k { + None => LogitsProcessor::new( + seed, + temperature.map(|temp| temp as f64), + top_p.map(|tp| tp as f64), + ), + Some(k) => { + let sampling = match temperature { + None => Sampling::ArgMax, + Some(temperature) => match top_p { + None => Sampling::TopK { + k, + temperature: temperature as f64, + }, + Some(p) => Sampling::TopKThenTopP { + k, + p: p as f64, + temperature: temperature as f64, + }, + }, + }; + LogitsProcessor::from_sampling(seed, sampling) + } + } +} + +pub fn use_repeat_penalty( + repeat_penalty: f32, + repeat_last_n: Option, + logits: &Tensor, + context: &[u32], +) -> Result { + if repeat_penalty == 1.0 || repeat_last_n.map_or(false, |n| n == 0) { + Ok(logits.clone()) + } else { + let start_at = if let Some(last_n) = repeat_last_n { + context.len().saturating_sub(last_n) + } else { + 0 + }; + Ok(candle_transformers::utils::apply_repeat_penalty( + &logits, + repeat_penalty, + &context[start_at..], + )?) + } +} + +/// logits shape: (dim) +pub fn simple_sample( + logits: &Tensor, + do_sample: bool, + temperature: Option, + top_k: Option, + top_p: Option, + previous_token_ids: Option<&[u32]>, + repeat_penalty: f32, + seed: Option, +) -> Result { + if logits.rank() != 1 { + return Err(anyhow!("simple_sample logits need rank = 1")); + } + let mut logits = if repeat_penalty != 1.0 + && let Some(tokens) = previous_token_ids + { + use_repeat_penalty(repeat_penalty, None, logits, tokens)? + } else { + logits.clone() + }; + if !do_sample { + Ok(logits.argmax(0)?.to_scalar::()?) + } else { + if let Some(temp) = temperature + && temp > 0.0 + { + logits = logits.affine(1.0 / temp, 0.0)?; + } + if let Some(top_k) = top_k + && top_k > 0 + && top_k > logits.dim(0)? + { + let sorted_indices = logits.arg_sort_last_dim(false)?; + let top_k_indices = sorted_indices.narrow(0, 0, top_k)?; + let top_k_logits = logits.gather(&top_k_indices, 0)?; + let threshold = top_k_logits.min_all()?; + let mask = logits.broadcast_lt(&threshold)?; + let on_true = Tensor::new(f32::NEG_INFINITY, logits.device())? + .to_dtype(logits.dtype())? + .broadcast_as(mask.shape())?; + logits = mask.where_cond(&on_true, &logits)?; + } + if let Some(top_p) = top_p + && top_p > 0.0 + && top_p < 1.0 + { + let sorted_indices = logits.arg_sort_last_dim(false)?; + let sorted_logits = logits.gather(&sorted_indices, 0)?; + let sorted_probs = softmax(&sorted_logits, 0)?; + let sorted_cumsum = sorted_probs.cumsum(0)?; + let mut mask = sorted_cumsum + .broadcast_gt(&Tensor::new(top_p, logits.device())?.to_dtype(logits.dtype())?)?; + // 保证数据不会被全部置为-inf + if mask.i(0)?.to_scalar::()? == 1 { + mask = mask.slice_scatter(&Tensor::new(0u32, logits.device())?, 0, 0)?; + } + let on_true = Tensor::new(f32::NEG_INFINITY, logits.device())? + .to_dtype(logits.dtype())? + .broadcast_as(mask.shape())?; + let new_logits = mask.where_cond(&on_true, &sorted_logits)?; + logits = logits.scatter(&sorted_indices, &new_logits, 0)?; + } + + let probs = softmax(&logits, 0)? + .to_dtype(candle_core::DType::F32)? + .to_vec1::()?; + let distr = rand::distr::weighted::WeightedIndex::new(probs).map_err(|e| { + anyhow!(format!( + "simple_sampel new rand::distr::weighted::WeightedIndex Failed: {}", + e + )) + })?; + let seed = seed.unwrap_or(34567); + let mut rng = rand::rngs::StdRng::seed_from_u64(seed); + let next_token = distr.sample(&mut rng) as u32; + Ok(next_token) + } +} diff --git a/src/models/fire_red_vad/processor.rs b/src/models/fire_red_vad/processor.rs index 4b72452..582ca6c 100644 --- a/src/models/fire_red_vad/processor.rs +++ b/src/models/fire_red_vad/processor.rs @@ -133,7 +133,8 @@ impl AudioFeat { } pub fn extract_file(&self, audio_path: &str, device: &Device) -> Result<(Tensor, f32)> { - let wave_tensor = load_audio_with_resample(audio_path, device, Some(16000))?.squeeze(0)?; + let wave_tensor = + load_audio_with_resample(audio_path, device, Some(16000), Some(1))?.squeeze(0)?; // fire_red_vad need i16 type data let wave_tensor = wave_tensor.affine(32768.0, 0.0)?; let dur = wave_tensor.dim(0)? as f32 / 16000.0; diff --git a/src/models/fire_red_vad/vad.rs b/src/models/fire_red_vad/vad.rs index 90acd18..6e02a5d 100644 --- a/src/models/fire_red_vad/vad.rs +++ b/src/models/fire_red_vad/vad.rs @@ -191,7 +191,7 @@ impl FireRedVad { return Err(anyhow!("only stream model support detect_frame")); } let audio_frame = - resample_audio_from_bytes(audio_bytes, &self.device, Some(16000))?.squeeze(0)?; + resample_audio_from_bytes(audio_bytes, &self.device, Some(16000), 1)?.squeeze(0)?; self.detect_frame(&audio_frame) } diff --git a/src/models/fun_asr_nano/processor.rs b/src/models/fun_asr_nano/processor.rs index 04adcff..372d6fb 100644 --- a/src/models/fun_asr_nano/processor.rs +++ b/src/models/fun_asr_nano/processor.rs @@ -91,7 +91,7 @@ impl FunAsrNanoProcessor { let sub_token = tokenizer.text_encode_vec(sub_prompt, true)?; source_ids.extend_from_slice(&sub_token); fbank_mask.extend_from_slice(&vec![0u32; sub_token.len()]); - let audio_tensors = extract_audios(mes, &self.device, Some(self.fronted_conf.fs))?; + let audio_tensors = extract_audios(mes, &self.device, Some(self.fronted_conf.fs), Some(1))?; if audio_tensors.is_empty() { return Err(anyhow!("FunASRNano need audio input")); } diff --git a/src/models/glm_asr_nano/model.rs b/src/models/glm_asr_nano/model.rs index e6d0a68..ab3f2c4 100644 --- a/src/models/glm_asr_nano/model.rs +++ b/src/models/glm_asr_nano/model.rs @@ -7,10 +7,10 @@ use crate::{ common::{ InferenceModel, modules::{ - LlamaForCausalLM, TwoLinearMLP, eager_attention_forward, get_conv1d, get_layer_norm, + TwoLinearMLP, eager_attention_forward, get_conv1d, get_layer_norm, }, }, - glm_asr_nano::config::{GlmAsrAudioConfig, GlmAsrNanoConfig}, + glm_asr_nano::config::{GlmAsrAudioConfig, GlmAsrNanoConfig}, llama::LlamaForCausalLM, }, position_embed::rope::{RoPE, glm_asr_apply_rotary_pos_emb}, utils::tensor_utils::{get_equal_mask, masked_scatter_dim0}, diff --git a/src/models/glm_asr_nano/processor.rs b/src/models/glm_asr_nano/processor.rs index 184e198..308569a 100644 --- a/src/models/glm_asr_nano/processor.rs +++ b/src/models/glm_asr_nano/processor.rs @@ -205,7 +205,7 @@ impl GlmAsrNanoProcessor { mes: &ChatCompletionParameters, render_text: &str, ) -> Result<(Tensor, Vec, String)> { - let audio_tensors = extract_audios(mes, &self.device, Some(self.sampling_rate))?; + let audio_tensors = extract_audios(mes, &self.device, Some(self.sampling_rate), Some(1))?; if audio_tensors.is_empty() { return Err(anyhow::anyhow!("GlmASRNano need audio input")); } diff --git a/src/models/gpt2/config.rs b/src/models/gpt2/config.rs new file mode 100644 index 0000000..998dd87 --- /dev/null +++ b/src/models/gpt2/config.rs @@ -0,0 +1,81 @@ +use serde::Deserialize; + +#[derive(Clone, Debug, Deserialize)] +pub struct GPT2Config { + pub activation_function: String, + pub add_cross_attention: bool, + pub attn_pdrop: f64, + pub bad_words_ids: Option>, + pub begin_suppress_tokens: Option>, + pub bos_token_id: u32, + pub chunk_size_feed_forward: usize, + pub cross_attention_hidden_size: Option, + pub decoder_start_token_id: Option, + pub diversity_penalty: f64, + pub do_sample: bool, + pub dtype: Option, + pub early_stopping: bool, + pub embd_pdrop: f64, + pub encoder_no_repeat_ngram_size: usize, + pub eos_token_id: u32, + pub exponential_decay_length_penalty: Option, + pub finetuning_task: Option, + pub forced_bos_token_id: Option, + pub forced_eos_token_id: Option, + pub id2label: std::collections::HashMap, + pub initializer_range: f64, + pub is_decoder: bool, + pub is_encoder_decoder: bool, + pub label2id: std::collections::HashMap, + pub layer_norm_epsilon: f64, + pub length_penalty: f64, + pub max_length: usize, + pub min_length: usize, + pub model_type: String, + pub n_ctx: usize, + pub n_embd: usize, + pub n_head: usize, + pub n_inner: usize, + pub n_layer: usize, + pub n_positions: usize, + pub no_repeat_ngram_size: usize, + pub num_beam_groups: usize, + pub num_beams: usize, + pub num_return_sequences: usize, + pub output_attentions: bool, + pub output_hidden_states: bool, + pub output_scores: bool, + pub pad_token_id: u32, + pub position_embedding_type: String, + pub prefix: Option, + pub problem_type: Option, + pub remove_invalid_values: bool, + pub reorder_and_upcast_attn: bool, + pub repetition_penalty: f64, + pub resid_pdrop: f64, + pub return_dict: bool, + pub return_dict_in_generate: bool, + pub rope_base: f64, + pub scale_attn_by_inverse_layer_idx: bool, + pub scale_attn_weights: bool, + pub sep_token_id: Option, + pub summary_activation: Option, + pub summary_first_dropout: f64, + pub summary_proj_to_labels: bool, + pub summary_type: String, + pub summary_use_proj: bool, + pub suppress_tokens: Option>, + pub task_specific_params: Option, + pub temperature: f64, + pub tf_legacy_loss: bool, + pub tie_encoder_decoder: bool, + pub tie_word_embeddings: bool, + pub tokenizer_class: Option, + pub top_k: usize, + pub top_p: f64, + pub torchscript: bool, + pub typical_p: f64, + pub use_bfloat16: bool, + pub use_cache: bool, + pub vocab_size: usize, +} diff --git a/src/models/gpt2/mod.rs b/src/models/gpt2/mod.rs new file mode 100644 index 0000000..261c5c5 --- /dev/null +++ b/src/models/gpt2/mod.rs @@ -0,0 +1,311 @@ +use anyhow::Result; +use candle_core::Tensor; +use candle_nn::{ + Activation, Embedding, Init, LayerNorm, Linear, Module, VarBuilder, embedding, linear_b, +}; + +use crate::{ + models::common::modules::{TwoLinearMLP, eager_attention_forward, get_layer_norm}, + position_embed::rope::{RoPE, apply_rotary_pos_emb_interleave}, + utils::tensor_utils::prepare_causal_attention_mask, +}; +pub mod config; + +pub struct GPT2Attention { + num_heads: usize, + head_dim: usize, + c_attn: Linear, + c_proj: Linear, + kv_cache: Option<(Tensor, Tensor)>, +} + +impl GPT2Attention { + pub fn new(vb: VarBuilder, hidden_size: usize, num_heads: usize) -> Result { + let c_attn = linear_b(hidden_size, 3 * hidden_size, true, vb.pp("c_attn"))?; + let c_proj = linear_b(hidden_size, hidden_size, true, vb.pp("c_proj"))?; + let head_dim = hidden_size / num_heads; + Ok(Self { + num_heads, + head_dim, + c_attn, + c_proj, + kv_cache: None, + }) + } + + pub fn forward( + &mut self, + xs: &Tensor, + cos: Option<&Tensor>, + sin: Option<&Tensor>, + attention_mask: Option<&Tensor>, + ) -> Result { + let (b, seq_len, _) = xs.dims3()?; + let xs = self.c_attn.forward(xs)?; + let xs_splits = xs.chunk(3, 2)?; + let query_states = xs_splits[0] + .as_ref() + .reshape((b, seq_len, self.num_heads, self.head_dim))? + .transpose(1, 2)?; + let key_states = xs_splits[1] + .as_ref() + .reshape((b, seq_len, self.num_heads, self.head_dim))? + .transpose(1, 2)?; + let value_states = xs_splits[2] + .as_ref() + .reshape((b, seq_len, self.num_heads, self.head_dim))? + .transpose(1, 2)?; + let (query_states, key_states) = if let Some(cos) = cos + && let Some(sin) = sin + { + apply_rotary_pos_emb_interleave(&query_states, &key_states, cos, sin, false)? + } else { + (query_states, key_states) + }; + let (key_states, value_states) = match &self.kv_cache { + None => (key_states, value_states), + Some((prev_k, prev_v)) => { + let key_states = Tensor::cat(&[prev_k, &key_states], 2)?; + let value_states = Tensor::cat(&[prev_v, &value_states], 2)?; + (key_states, value_states) + } + }; + self.kv_cache = Some((key_states.clone(), value_states.clone())); + let scale = 1f64 / f64::sqrt(self.head_dim as f64); + let attn_output = eager_attention_forward( + &query_states, + &key_states, + &value_states, + None, + attention_mask, + scale, + )?; + let attn_output = attn_output.reshape((b, seq_len, self.num_heads * self.head_dim))?; + let attn_output = attn_output.apply(&self.c_proj)?; + Ok(attn_output) + } + + pub fn clear_kv_cache(&mut self) { + self.kv_cache = None + } +} + +pub struct GPT2MLP { + linear1: Linear, + linear2: Linear, + act: Activation, +} + +impl GPT2MLP { + pub fn new( + vb: VarBuilder, + in_dim: usize, + middle_dim: usize, + out_dim: usize, + act: Activation, + ) -> Result { + let c_fc_weight = vb + .get_with_hints((in_dim, middle_dim), "c_fc.weight", Init::Const(1.0))? + .t()?; + let c_fc_bias = vb.get_with_hints(middle_dim, "c_fc.bias", Init::Const(0.0))?; + let c_fc = Linear::new(c_fc_weight, Some(c_fc_bias)); + + let c_proj_weight = vb + .get_with_hints((middle_dim, out_dim), "c_proj.weight", Init::Const(1.0))? + .t()?; + let c_proj_bias = vb.get_with_hints(out_dim, "c_proj.bias", Init::Const(0.0))?; + let c_proj = Linear::new(c_proj_weight, Some(c_proj_bias)); + + Ok(Self { + linear1: c_fc, + linear2: c_proj, + act, + }) + } + pub fn forward(&self, xs: &Tensor) -> Result { + let xs = xs + .apply(&self.linear1)? + .apply(&self.act)? + .apply(&self.linear2)?; + Ok(xs) + } +} + +pub struct GPT2Block { + ln_1: LayerNorm, + attn: GPT2Attention, + ln_2: LayerNorm, + mlp: TwoLinearMLP, +} + +impl GPT2Block { + pub fn new( + vb: VarBuilder, + hidden_size: usize, + num_heads: usize, + inner_dim: Option, + ) -> Result { + let inner_dim = inner_dim.unwrap_or(4 * hidden_size); + let ln_1 = get_layer_norm(vb.pp("ln_1"), 1e-5, hidden_size, true)?; + let attn = GPT2Attention::new(vb.pp("attn"), hidden_size, num_heads)?; + let ln_2 = get_layer_norm(vb.pp("ln_2"), 1e-5, hidden_size, true)?; + let mlp = TwoLinearMLP::new( + vb.pp("mlp"), + hidden_size, + inner_dim, + hidden_size, + Activation::NewGelu, + true, + "fc_in", + "fc_out", + )?; + Ok(Self { + ln_1, + attn, + ln_2, + mlp, + }) + } + + pub fn forward( + &mut self, + xs: &Tensor, + cos: Option<&Tensor>, + sin: Option<&Tensor>, + attention_mask: Option<&Tensor>, + ) -> Result { + let residual = xs.clone(); + let xs = self.ln_1.forward(xs)?; + let xs = self.attn.forward(&xs, cos, sin, attention_mask)?; + let residual = xs.add(&residual)?; + let xs = self.ln_2.forward(&residual)?; + let xs = self.mlp.forward(&xs)?; + let xs = xs.add(&residual)?; + Ok(xs) + } + pub fn clear_kv_cache(&mut self) { + self.attn.clear_kv_cache() + } +} + +#[allow(unused)] +pub struct GPT2Model { + pub wte: Option, + // wpe: Embedding, //rope not need + h: Vec, + ln_f: LayerNorm, + rope: Option, +} + +#[allow(unused)] +impl GPT2Model { + pub fn new( + vb: VarBuilder, + hidden_size: usize, + num_heads: usize, + num_hidden_layers: usize, + vocab_size: usize, + // n_positions: usize, + ) -> Result { + let wte = Some(embedding(vocab_size, hidden_size, vb.pp("wte"))?); + let vb_layers = vb.pp("h"); + let mut h = vec![]; + for i in 0..num_hidden_layers { + let block = GPT2Block::new(vb_layers.pp(i), hidden_size, num_heads, None)?; + h.push(block); + } + let ln_f = get_layer_norm(vb.pp("ln_f"), 1e-5, hidden_size, true)?; + let head_dim = hidden_size / num_heads; + let rope = RoPE::new(head_dim, 10000.0, vb.device())?; + Ok(Self { + wte, + h, + ln_f, + rope: Some(rope), + }) + } + + pub fn new_without_wte( + vb: VarBuilder, + hidden_size: usize, + num_heads: usize, + num_hidden_layers: usize, + vocab_size: usize, + // n_positions: usize, + ) -> Result { + let wte = None; + let vb_layers = vb.pp("h"); + let mut h = vec![]; + for i in 0..num_hidden_layers { + let block = GPT2Block::new(vb_layers.pp(i), hidden_size, num_heads, None)?; + h.push(block); + } + let ln_f = get_layer_norm(vb.pp("ln_f"), 1e-5, hidden_size, true)?; + let head_dim = hidden_size / num_heads; + let rope = RoPE::new(head_dim, 10000.0, vb.device())?; + Ok(Self { + wte, + h, + ln_f, + rope: Some(rope), + }) + } + + pub fn new_with_wte( + vb: VarBuilder, + hidden_size: usize, + num_heads: usize, + num_hidden_layers: usize, + wte_embeddings: &Tensor, + ) -> Result { + let wte = Some(Embedding::new(wte_embeddings.clone(), hidden_size)); + let vb_layers = vb.pp("h"); + let mut h = vec![]; + for i in 0..num_hidden_layers { + let block = GPT2Block::new(vb_layers.pp(i), hidden_size, num_heads, None)?; + h.push(block); + } + let ln_f = get_layer_norm(vb.pp("ln_f"), 1e-5, hidden_size, true)?; + Ok(Self { + wte, + h, + ln_f, + rope: None, + }) + } + + pub fn forward(&mut self, inputs_embeds: &Tensor, seqlen_offset: usize) -> Result { + let (b_size, seq_len, _) = inputs_embeds.dims3()?; + let (cos, sin) = if let Some(rope) = &self.rope { + let (cos, sin) = rope.forward_repeat_interleave(seqlen_offset, seq_len, inputs_embeds.device())?; + (Some(cos), Some(sin)) + } else { + (None, None) + }; + + let mut xs = inputs_embeds.clone(); + let attention_mask: Option = { + if seq_len <= 1 { + None + } else { + Some(prepare_causal_attention_mask( + b_size, + seq_len, + 0, + xs.device(), + )?) + } + }; + for block in &mut self.h { + xs = block.forward(&xs, cos.as_ref(), sin.as_ref(), attention_mask.as_ref())?; + } + xs = self.ln_f.forward(&xs)?; + Ok(xs) + } + + pub fn clear_kv_cache(&mut self) { + for layer in self.h.iter_mut() { + layer.clear_kv_cache() + } + } +} diff --git a/src/models/llama/mod.rs b/src/models/llama/mod.rs new file mode 100644 index 0000000..dff27e0 --- /dev/null +++ b/src/models/llama/mod.rs @@ -0,0 +1,166 @@ +use crate::{ + models::common::modules::NaiveAttnGateUpDownMLPBlock, position_embed::rope::RoPE, + utils::tensor_utils::prepare_causal_attention_mask, +}; +use anyhow::Result; +use candle_core::Tensor; +use candle_nn::{ + Activation, Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, linear_no_bias, rms_norm, +}; + +pub struct LlamaModel { + pub embed_tokens: Embedding, + layers: Vec, + norm: RmsNorm, + rotary_emb: RoPE, +} + +impl LlamaModel { + pub fn new( + vb: VarBuilder, + vocab_size: usize, + hidden_size: usize, + num_hidden_layers: usize, + num_attention_heads: usize, + num_key_value_heads: Option, + head_dim: Option, + attn_bias: bool, + attn_pp_name: &str, + o_proj_pp_name: Option<&str>, + intermediate_size: usize, + hidden_act: Activation, + mlp_bias: bool, + mlp_pp_name: &str, + norm_eps: f64, + input_norm_pp_name: &str, + post_norm_pp_name: &str, + rope_theta_base: f32, + ) -> Result { + let embed_tokens = embedding(vocab_size, hidden_size, vb.pp("embed_tokens"))?; + let mut layers = vec![]; + let vb_layers = vb.pp("layers"); + for i in 0..num_hidden_layers { + let layers_i = NaiveAttnGateUpDownMLPBlock::new( + vb_layers.pp(i), + hidden_size, + num_attention_heads, + num_key_value_heads, + head_dim, + attn_bias, + attn_pp_name, + o_proj_pp_name, + intermediate_size, + hidden_act, + mlp_bias, + mlp_pp_name, + norm_eps, + input_norm_pp_name, + post_norm_pp_name, + )?; + layers.push(layers_i); + } + let norm = rms_norm(hidden_size, norm_eps, vb.pp("norm"))?; + let head_dim = head_dim.unwrap_or(hidden_size / num_attention_heads); + let rotary_emb = RoPE::new(head_dim, rope_theta_base, vb.device())?; + Ok(Self { + embed_tokens, + layers, + norm, + rotary_emb, + }) + } + + pub fn forward(&mut self, inputs_embeds: &Tensor, seqlen_offset: usize) -> Result { + let (b_size, seq_len, _) = inputs_embeds.dims3()?; + + let (cos, sin) = self + .rotary_emb + .forward(seqlen_offset, seq_len, inputs_embeds.device())?; + let mut xs = inputs_embeds.clone(); + let attention_mask: Option = { + if seq_len <= 1 { + None + } else { + Some(prepare_causal_attention_mask( + b_size, + seq_len, + 0, + xs.device(), + )?) + } + }; + for layer in self.layers.iter_mut() { + xs = layer.forward(&xs, &cos, &sin, attention_mask.as_ref())?; + } + let xs = xs.apply(&self.norm)?; + Ok(xs) + } + + pub fn clear_kv_cache(&mut self) { + for layer in self.layers.iter_mut() { + layer.clear_kv_cache() + } + } +} + +pub struct LlamaForCausalLM { + pub model: LlamaModel, + lm_head: Linear, +} + +impl LlamaForCausalLM { + pub fn new( + vb: VarBuilder, + vocab_size: usize, + hidden_size: usize, + num_hidden_layers: usize, + num_attention_heads: usize, + num_key_value_heads: Option, + head_dim: Option, + attn_bias: bool, + attn_pp_name: &str, + o_proj_pp_name: Option<&str>, + intermediate_size: usize, + hidden_act: Activation, + mlp_bias: bool, + mlp_pp_name: &str, + norm_eps: f64, + input_norm_pp_name: &str, + post_norm_pp_name: &str, + rope_theta_base: f32, + ) -> Result { + let model = LlamaModel::new( + vb.pp("model"), + vocab_size, + hidden_size, + num_hidden_layers, + num_attention_heads, + num_key_value_heads, + head_dim, + attn_bias, + attn_pp_name, + o_proj_pp_name, + intermediate_size, + hidden_act, + mlp_bias, + mlp_pp_name, + norm_eps, + input_norm_pp_name, + post_norm_pp_name, + rope_theta_base, + )?; + let lm_head = linear_no_bias(hidden_size, vocab_size, vb.pp("lm_head"))?; + Ok(Self { model, lm_head }) + } + + pub fn forward(&mut self, inputs_embeds: &Tensor, seqlen_offset: usize) -> Result { + let outputs = self.model.forward(inputs_embeds, seqlen_offset)?; + let seq_len = outputs.dim(1)?; + let hidden_state = outputs.narrow(1, seq_len - 1, 1)?; + let logits = self.lm_head.forward(&hidden_state)?; + Ok(logits) + } + pub fn clear_kv_cache(&mut self) { + self.model.clear_kv_cache(); + } +} diff --git a/src/models/mask_gct/model.rs b/src/models/mask_gct/model.rs index bd2bb04..24389e8 100644 --- a/src/models/mask_gct/model.rs +++ b/src/models/mask_gct/model.rs @@ -128,6 +128,8 @@ impl FactorizedVectorQuantize { 1, 1, true, + None, + None, )?; let out_project = WNConv1d::new( vb.pp("out_project"), @@ -139,6 +141,8 @@ impl FactorizedVectorQuantize { 1, 1, true, + None, + None, )?; (Some(in_project), Some(out_project)) } else { diff --git a/src/models/mod.rs b/src/models/mod.rs index 6cb658c..64d369e 100644 --- a/src/models/mod.rs +++ b/src/models/mod.rs @@ -12,6 +12,7 @@ pub mod lfm2; pub mod lfm2vl; pub mod mask_gct; pub mod minicpm4; +pub mod moss; pub mod paddleocr_vl; pub mod qwen2; pub mod qwen2_5vl; @@ -27,6 +28,8 @@ pub mod voxcpm_refact; pub mod w2v_bert_2_0; // pub mod sam3; pub mod fire_red_vad; +pub mod gpt2; +pub mod llama; use crate::{ models::{ diff --git a/src/models/moss/audio_tokenizer_nano.rs b/src/models/moss/audio_tokenizer_nano.rs new file mode 100644 index 0000000..2c814ee --- /dev/null +++ b/src/models/moss/audio_tokenizer_nano.rs @@ -0,0 +1,669 @@ +use anyhow::{Result, anyhow}; +use candle_core::{D, IndexOp, Tensor}; +use candle_nn::{Embedding, LayerNorm, Linear, Module, VarBuilder, embedding, linear_no_bias}; + +use crate::{ + models::{ + common::modules::{ + TwoLinearMLP, WNConv1d, eager_attention_forward, get_layer_norm, l2_normalize, + }, + moss::config::{ + MossAudioTokenizerConfig, MossAudioTokenizerModuleConfig, + MossAudioTokenizerQuantizerKwargs, + }, + }, + position_embed::rope::{RoPE, apply_rotary_pos_emb_roformer}, +}; + +pub struct MossAudioTokenizerPatchedPretransform { + patch_size: usize, + is_downsample: bool, +} + +impl MossAudioTokenizerPatchedPretransform { + pub fn new(patch_size: usize, is_downsample: bool) -> Self { + Self { + patch_size, + is_downsample, + } + } + + pub fn encode(&self, x: &Tensor, input_lengths: &Tensor) -> Result<(Tensor, Tensor)> { + let (b, d, _) = x.dims3()?; + let x = x + .reshape((b, d, (), self.patch_size))? + .permute((0, 1, 3, 2))? + .reshape((b, d * self.patch_size, ()))?; + let out_lengths = input_lengths + .affine(1.0 / self.patch_size as f64, 0.0)? + .floor()?; + Ok((x, out_lengths)) + } + + pub fn decode(&self, x: &Tensor, input_lengths: &Tensor) -> Result<(Tensor, Tensor)> { + let (b, dh, l) = x.dims3()?; + let d = dh / self.patch_size; + let x = x + .reshape((b, d, self.patch_size, l))? + .permute((0, 1, 3, 2))? + .reshape((b, d, l * self.patch_size))?; + let out_lengths = (input_lengths * self.patch_size as f64)?; + Ok((x, out_lengths)) + } + + pub fn forward(&self, x: &Tensor, input_lengths: &Tensor) -> Result<(Tensor, Tensor)> { + if self.is_downsample { + self.encode(x, input_lengths) + } else { + self.decode(x, input_lengths) + } + } +} + +pub struct MossAudioTokenizerMultiheadAttention { + num_heads: usize, + scale: f64, + in_proj: Linear, + out_proj: Linear, +} + +impl MossAudioTokenizerMultiheadAttention { + pub fn new(vb: VarBuilder, embed_dim: usize, num_heads: usize) -> Result { + let head_dim = embed_dim / num_heads; + let scale = 1f64 / f64::sqrt(head_dim as f64); + let in_proj = linear_no_bias(embed_dim, 3 * embed_dim, vb.pp("in_proj"))?; + let out_proj = linear_no_bias(embed_dim, embed_dim, vb.pp("out_proj"))?; + Ok(Self { + num_heads, + scale, + in_proj, + out_proj, + }) + } + + pub fn forward( + &self, + xs: &Tensor, + cos: &Tensor, + sin: &Tensor, + mask: &Tensor, + input_lengths: &Tensor, + ) -> Result { + let (bs, max_seqlen, _) = xs.dims3()?; + let projected = self + .in_proj + .forward(xs)? + .reshape((bs, max_seqlen, 3, self.num_heads, ()))? + .permute((2, 0, 3, 1, 4))?; + let [q, k, v] = projected + .chunk(3, 0)? + .try_into() + .map_err(|_| anyhow!("Chunk size mismatch"))?; + let q = q.squeeze(0)?.contiguous()?; + let k = k.squeeze(0)?.contiguous()?; + let v = v.squeeze(0)?.contiguous()?; + let (q, k) = apply_rotary_pos_emb_roformer(&q, &k, cos, sin)?; + // let (q, k) = self.apply_rope(&q, &k, cos, sin)?; + let attn = eager_attention_forward(&q, &k, &v, None, Some(mask), self.scale)?; + // (b, seq_len, n_head, dim) -> (b, n_head, seq_len, dim) + let attn = attn.transpose(1, 2)?; + let valid_q = Tensor::arange(0f32, max_seqlen as f32, xs.device())? + .reshape((1, 1, max_seqlen, 1))? + .broadcast_lt( + &input_lengths + .reshape((bs, 1, 1, 1))? + .repeat((1, 1, max_seqlen, 1))?, + )? + .broadcast_as(attn.shape())?; + let on_false = attn.zeros_like()?; + let attn = valid_q.where_cond(&attn, &on_false)?; + // (b, n_head, seq_len, dim) -> (b, seq_len, n_head, dim) + let attn = attn.transpose(1, 2)?; + let attn = attn.reshape((bs, max_seqlen, ()))?; + let out = self.out_proj.forward(&attn)?; + Ok(out) + } +} + +pub struct MossAudioTokenizerTransformerLayer { + self_attn: MossAudioTokenizerMultiheadAttention, + norm1: LayerNorm, + norm2: LayerNorm, + ffn: TwoLinearMLP, + layer_scale_1: Tensor, + layer_scale_2: Tensor, +} + +impl MossAudioTokenizerTransformerLayer { + pub fn new(vb: VarBuilder, config: &MossAudioTokenizerModuleConfig) -> Result { + let self_attn = MossAudioTokenizerMultiheadAttention::new( + vb.pp("self_attn"), + config.d_model.unwrap(), + config.num_heads.unwrap(), + )?; + let norm1 = get_layer_norm(vb.pp("norm1"), 1e-5, config.d_model.unwrap(), true)?; + let norm2 = get_layer_norm(vb.pp("norm2"), 1e-5, config.d_model.unwrap(), true)?; + let ffn = TwoLinearMLP::new( + vb.pp("ffn"), + config.d_model.unwrap(), + config.dim_feedforward.unwrap(), + config.d_model.unwrap(), + candle_nn::Activation::Gelu, + false, + "0", + "2", + )?; + let layer_scale_1 = vb + .get(config.d_model.unwrap(), "layer_scale_1.scale")? + .unsqueeze(0)? + .unsqueeze(0)?; + let layer_scale_2 = vb + .get(config.d_model.unwrap(), "layer_scale_2.scale")? + .unsqueeze(0)? + .unsqueeze(0)?; + Ok(Self { + self_attn, + norm1, + norm2, + ffn, + layer_scale_1, + layer_scale_2, + }) + } + + pub fn forward( + &self, + xs: &Tensor, + cos: &Tensor, + sin: &Tensor, + mask: &Tensor, + input_lengths: &Tensor, + ) -> Result { + let residual = xs.clone(); + let xs = self.norm1.forward(xs)?; + let xs = self.self_attn.forward(&xs, cos, sin, mask, input_lengths)?; + let xs = self.layer_scale_1.broadcast_mul(&xs)?; + let residual = residual.add(&xs)?; + let xs = self.norm2.forward(&residual)?; + let xs = self.ffn.forward(&xs)?; + let xs = self.layer_scale_2.broadcast_mul(&xs)?; + let xs = residual.add(&xs)?; + Ok(xs) + } +} + +pub struct MossAudioTokenizerTransformer { + rope: RoPE, // use roformer + context: usize, + layers: Vec, +} + +impl MossAudioTokenizerTransformer { + pub fn new( + vb: VarBuilder, + config: &MossAudioTokenizerModuleConfig, + context: usize, + ) -> Result { + let dim = config.d_model.unwrap() / config.num_heads.unwrap(); + let rope = RoPE::new(dim, 10000.0, vb.device())?; + let vb_layers = vb.pp("layers"); + let mut layers = vec![]; + for i in 0..config.num_layers.unwrap() { + let layer = MossAudioTokenizerTransformerLayer::new(vb_layers.pp(i), config)?; + layers.push(layer); + } + Ok(Self { + rope, + context, + layers, + }) + } + + pub fn forward(&self, input_embeds: &Tensor, input_lengths: &Tensor) -> Result { + let t = input_embeds.dim(1)?; + let (cos, sin) = self.rope.forward(0, t, input_embeds.device())?; + let mut xs = input_embeds.clone(); + let mask = self.build_attn_bias(input_lengths, t)?; + for layer in self.layers.iter() { + xs = layer.forward(&xs, &cos, &sin, &mask, input_lengths)?; + } + Ok(xs) + } + + fn build_attn_bias(&self, input_lengths: &Tensor, max_seqlen: usize) -> Result { + let positions = Tensor::arange(0f32, max_seqlen as f32, input_lengths.device())?; + let input_lengths = input_lengths.reshape(((), 1, 1))?; + let valid_k = positions + .reshape((1, 1, max_seqlen))? + .broadcast_lt(&input_lengths)?; + let delta = positions + .reshape((1, max_seqlen, 1))? + .broadcast_sub(&positions.reshape((1, 1, max_seqlen))?)?; + let delta1 = delta.ge(&delta.zeros_like()?)?; + let delta2 = delta + .lt(&Tensor::new(self.context as f32, delta.device())?.broadcast_as(delta.shape())?)?; + let mask = delta1.minimum(&delta2)?; + let mask = mask.broadcast_minimum(&valid_k)?.unsqueeze(1)?; + let on_true = mask.zeros_like()?.to_dtype(candle_core::DType::F32)?; + let on_false = Tensor::new(f32::NEG_INFINITY, mask.device())?.broadcast_as(mask.shape())?; + let mask = mask.where_cond(&on_true, &on_false)?; + Ok(mask) + } +} + +pub struct MossAudioTokenizerProjectedTransformer { + input_proj: Linear, + transformer: MossAudioTokenizerTransformer, + output_proj: Linear, +} + +impl MossAudioTokenizerProjectedTransformer { + pub fn new( + vb: VarBuilder, + config: &MossAudioTokenizerModuleConfig, + context: usize, + ) -> Result { + let input_proj = linear_no_bias( + config.input_dimension.unwrap(), + config.d_model.unwrap(), + vb.pp("input_proj"), + )?; + let transformer = + MossAudioTokenizerTransformer::new(vb.pp("transformer"), config, context)?; + let output_proj = linear_no_bias( + config.d_model.unwrap(), + config.output_dimension.unwrap(), + vb.pp("output_proj"), + )?; + Ok(Self { + input_proj, + transformer, + output_proj, + }) + } + + pub fn forward( + &self, + input_embeds: &Tensor, + input_lengths: &Tensor, + ) -> Result<(Tensor, Tensor)> { + let xs = self.input_proj.forward(&input_embeds.transpose(1, 2)?)?; + let xs = self.transformer.forward(&xs, input_lengths)?; + let xs = self.output_proj.forward(&xs)?.transpose(1, 2)?; + Ok((xs, input_lengths.clone())) + } +} + +pub enum MossAudioTokenizerModule { + PatchedPretransform(MossAudioTokenizerPatchedPretransform), + ProjectedTransformer(MossAudioTokenizerProjectedTransformer), +} + +impl MossAudioTokenizerModule { + pub fn forward( + &self, + input_embeds: &Tensor, + input_lengths: &Tensor, + ) -> Result<(Tensor, Tensor)> { + match self { + MossAudioTokenizerModule::PatchedPretransform(patch) => { + patch.forward(input_embeds, input_lengths) + } + MossAudioTokenizerModule::ProjectedTransformer(transformer) => { + transformer.forward(input_embeds, input_lengths) + } + } + } +} + +pub struct MossAudioTokenizerLFQ { + in_proj: Option, + out_proj: Option, + codebook: Embedding, + codebook_l2_norm: Tensor, +} + +impl MossAudioTokenizerLFQ { + pub fn new(vb: VarBuilder, config: &MossAudioTokenizerQuantizerKwargs) -> Result { + let in_proj = if config.rvq_dim != config.codebook_dim { + Some(WNConv1d::new( + vb.pp("in_proj"), + config.rvq_dim, + config.codebook_dim, + 1, + 1, + 0, + 1, + 1, + true, + Some("parametrizations.weight.original0"), + Some("parametrizations.weight.original1"), + )?) + } else { + None + }; + + let out_proj = if config.rvq_dim != config.codebook_dim { + Some(WNConv1d::new( + vb.pp("out_proj"), + config.codebook_dim, + config.rvq_dim, + 1, + 1, + 0, + 1, + 1, + true, + Some("parametrizations.weight.original0"), + Some("parametrizations.weight.original1"), + )?) + } else { + None + }; + + let codebook = embedding(config.codebook_size, config.codebook_dim, vb.pp("codebook"))?; + let codebook_l2_norm = l2_normalize(codebook.embeddings(), 1)?; + Ok(Self { + in_proj, + out_proj, + codebook, + codebook_l2_norm, + }) + } + + pub fn forward(&self, xs: &Tensor) -> Result<(Tensor, Tensor)> { + let z_e = if let Some(in_proj) = &self.in_proj { + in_proj.forward(xs)? + } else { + xs.clone() + }; + let (bs, len, _) = z_e.dims3()?; + let encodings = z_e.transpose(1, 2)?.reshape(((), len))?; + let encodings = l2_normalize(&encodings, 1)?; + let dist1 = encodings.powf(2.0)?.sum_keepdim(1)?; + let dist2 = encodings + .affine(2.0, 0.0)? + .matmul(&self.codebook_l2_norm.t()?)?; + let dist3 = self.codebook_l2_norm.powf(2.0)?.sum_keepdim(1)?.t()?; + let dist = dist1.broadcast_sub(&dist2)?.broadcast_add(&dist3)?; + let indices = dist + .affine(-1.0, 0.0)? + .argmax(1)? + .reshape((bs, ()))? + .to_dtype(candle_core::DType::U32)?; + let z_q = self.codebook.forward(&indices)?.transpose(1, 2)?; + let mut z_q = z_e.add(&z_q.sub(&z_e)?)?; + if let Some(out_proj) = &self.out_proj { + z_q = out_proj.forward(&z_q)?; + } + Ok((z_q, indices)) + } +} + +pub struct MossAudioTokenizerResidualLFQ { + input_proj: Option, + output_proj: Option, + quantizers: Vec, +} + +impl MossAudioTokenizerResidualLFQ { + pub fn new(vb: VarBuilder, config: &MossAudioTokenizerQuantizerKwargs) -> Result { + let input_proj = if config.input_dim != config.rvq_dim { + Some(WNConv1d::new( + vb.pp("input_proj"), + config.input_dim, + config.rvq_dim, + 1, + 1, + 0, + 1, + 1, + true, + Some("parametrizations.weight.original0"), + Some("parametrizations.weight.original1"), + )?) + } else { + None + }; + + let output_proj = if config.rvq_dim != config.output_dim { + Some(WNConv1d::new( + vb.pp("output_proj"), + config.rvq_dim, + config.output_dim, + 1, + 1, + 0, + 1, + 1, + true, + Some("parametrizations.weight.original0"), + Some("parametrizations.weight.original1"), + )?) + } else { + None + }; + + let vb_quantizers = vb.pp("quantizers"); + let mut quantizers = vec![]; + for i in 0..config.num_quantizers { + let layer = MossAudioTokenizerLFQ::new(vb_quantizers.pp(i), config)?; + quantizers.push(layer); + } + Ok(Self { + input_proj, + output_proj, + quantizers, + }) + } + + pub fn forward(&self, input_values: &Tensor, length: &Tensor) -> Result { + let z = if let Some(proj) = &self.input_proj { + proj.forward(input_values)? + } else { + input_values.clone() + }; + let max_time = z.dim(2)?; + let mask = Tensor::arange(0f32, max_time as f32, z.device())? + .unsqueeze(0)? + .broadcast_lt(&length.unsqueeze(1)?)? + .unsqueeze(1)?; + // let mut quantized_out = z.zeros_like()?; + let mut residual = z.clone(); + let on_false = residual.zeros_like()?; + let mask_reshape = mask.broadcast_as(residual.shape())?; + let mut all_indices = vec![]; + for quantizer in self.quantizers.iter() { + let masked_residual = mask_reshape.where_cond(&residual, &on_false)?; + let (z_q_i, indices_i) = quantizer.forward(&masked_residual)?; + all_indices.push(indices_i); + let z_q_i_mask = mask_reshape.where_cond(&z_q_i, &on_false)?; + residual = residual.sub(&z_q_i_mask)?; + } + let all_indices = Tensor::stack(&all_indices, 0)?; + Ok(all_indices) + } +} + +pub struct MossAudioTokenizer { + pub sampling_rate: usize, + pub downsample_rate: usize, + pub number_channels: usize, + pub enable_channel_interleave: bool, + encoder: Vec, + quantizer: MossAudioTokenizerResidualLFQ, + decoder: Vec, +} + +impl MossAudioTokenizer { + pub fn new(vb: VarBuilder, config: &MossAudioTokenizerConfig) -> Result { + let channel_interleave_factor = + if config.enable_channel_interleave && config.number_channels > 1 { + config.number_channels + } else { + 1 + }; + let current_frame_rate = config.sampling_rate * channel_interleave_factor; + let vb_encoder = vb.pp("encoder"); + let mut encoder = vec![]; + for (layer_id, cfg) in config.encoder_kwargs.iter().enumerate() { + if cfg.module_type == "PatchedPretransform" + && let Some(patch_size) = cfg.patch_size + { + let layer = MossAudioTokenizerPatchedPretransform::new(patch_size, true); + encoder.push(MossAudioTokenizerModule::PatchedPretransform(layer)); + } else if cfg.module_type == "Transformer" { + let context_duration = cfg + .context_duration + .unwrap_or(config.causal_transformer_context_duration); + let context = (current_frame_rate as f64 * context_duration).round() as usize; + let layer = MossAudioTokenizerProjectedTransformer::new( + vb_encoder.pp(layer_id), + cfg, + context, + )?; + encoder.push(MossAudioTokenizerModule::ProjectedTransformer(layer)); + } else { + return Err(anyhow!( + "Moss Module only sopport PatchedPretransform and Transformer, but get: {}", + cfg.module_type + )); + } + } + let quantizer = + MossAudioTokenizerResidualLFQ::new(vb.pp("quantizer"), &config.quantizer_kwargs)?; + let vb_decoder = vb.pp("decoder"); + let mut decoder = vec![]; + for (layer_id, cfg) in config.decoder_kwargs.iter().enumerate() { + if cfg.module_type == "PatchedPretransform" + && let Some(patch_size) = cfg.patch_size + { + let layer = MossAudioTokenizerPatchedPretransform::new(patch_size, true); + decoder.push(MossAudioTokenizerModule::PatchedPretransform(layer)); + } else if cfg.module_type == "Transformer" { + let context_duration = cfg + .context_duration + .unwrap_or(config.causal_transformer_context_duration); + let context = (current_frame_rate as f64 * context_duration).round() as usize; + let layer = MossAudioTokenizerProjectedTransformer::new( + vb_decoder.pp(layer_id), + cfg, + context, + )?; + decoder.push(MossAudioTokenizerModule::ProjectedTransformer(layer)); + } else { + return Err(anyhow!( + "Moss Module only sopport PatchedPretransform and Transformer, but get: {}", + cfg.module_type + )); + } + } + + Ok(Self { + sampling_rate: config.sampling_rate, + downsample_rate: config.downsample_rate, + number_channels: config.number_channels, + enable_channel_interleave: config.enable_channel_interleave, + encoder, + quantizer, + decoder, + }) + } + + fn flatten_channels_for_codec( + &self, + input_values: &Tensor, + length: &Tensor, + ) -> Result<(Tensor, Tensor)> { + let (bs, _, audio_len) = input_values.dims3()?; + let input_values = if audio_len % self.downsample_rate != 0 { + let pad_length = self.downsample_rate - (audio_len % self.downsample_rate); + input_values.pad_with_zeros(D::Minus1, 0, pad_length)? + } else { + input_values.clone() + }; + if self.number_channels > 1 && self.enable_channel_interleave { + let input_values = input_values + .transpose(1, 2)? + .contiguous()? + .reshape((bs, 1, ()))?; + let length = (length * self.number_channels as f64)?; + Ok((input_values, length)) + } else { + Ok((input_values, length.clone())) + } + } + + pub fn batch_encode(&self, input_values: &Tensor, length: &Tensor) -> Result> { + let (mut encoder_hidden_states, mut encoder_hidden_lengths) = + self.flatten_channels_for_codec(input_values, length)?; + for layer in &self.encoder { + (encoder_hidden_states, encoder_hidden_lengths) = + layer.forward(&encoder_hidden_states, &encoder_hidden_lengths)?; + } + let audio_codes = self + .quantizer + .forward(&encoder_hidden_states, &encoder_hidden_lengths)?; + // (dim, bs, len) -> (bs, len, dim) + let audio_codes = audio_codes.permute((1, 2, 0))?; + let mut audio_codes_vec = vec![]; + for index in 0..encoder_hidden_lengths.dim(0)? { + let codes_i = audio_codes.i(index)?; + let length = encoder_hidden_lengths.i(index)?.to_scalar::()? as usize; + let codes_i = codes_i.narrow(0, 0, length)?; + audio_codes_vec.push(codes_i); + } + Ok(audio_codes_vec) + } + + pub fn encode_one(&self, wav: &Tensor) -> Result { + // (channel, audio_len) -> (bs=1, channel, audio_len) + let (c, len) = wav.dims2()?; + if c != self.number_channels { + return Err(anyhow!( + "MossAudioTokenizer encode_one need number_channels: {} but the wav channel: {}", + self.number_channels, + c, + )); + } + let input_values = wav.unsqueeze(0)?; + let length = Tensor::new(vec![len as f32], wav.device())?; + let audio_vec = self.batch_encode(&input_values, &length)?; + Ok(audio_vec[0].clone()) + } + + pub fn encode_list(&self, wavs: &Vec) -> Result> { + if wavs.is_empty() { + return Err(anyhow!( + "MossAudioTokenizer encode_list need wavs len > 0, but the wavs is empty" + )); + } + let mut length = vec![]; + for wav in wavs.iter() { + let (c, len) = wav.dims2()?; + if c != self.number_channels { + return Err(anyhow!( + "MossAudioTokenizer encode_list need number_channels: {} but the wav channel: {}", + self.number_channels, + c, + )); + } + length.push(len as u32); + } + let max_length = *length.iter().max().unwrap_or(&0) as usize; + let mut input_values = vec![]; + for wav in wavs.iter() { + let audio_len = wav.dim(1)?; + let wav_ = if audio_len < max_length { + wav.pad_with_zeros(D::Minus1, 0, max_length - audio_len)? + } else { + wav.clone() + }; + input_values.push(wav_); + } + let input_values = Tensor::stack(&input_values, 0)?; + let length_tensor = Tensor::new(length.clone(), input_values.device())? + .to_dtype(candle_core::DType::F32)?; + Ok(self.batch_encode(&input_values, &length_tensor)?) + } +} diff --git a/src/models/moss/config.rs b/src/models/moss/config.rs new file mode 100644 index 0000000..197e163 --- /dev/null +++ b/src/models/moss/config.rs @@ -0,0 +1,159 @@ +use serde::Deserialize; + +use crate::models::gpt2::config::GPT2Config; + +#[derive(Debug, Deserialize)] +pub struct MossAudioTokenizerConfig { + pub sample_rate: usize, + pub sampling_rate: usize, + pub downsample_rate: usize, + pub causal_transformer_context_duration: f64, + pub number_channels: usize, + pub enable_channel_interleave: bool, + pub compute_dtype: String, + pub dtype: String, + pub code_dim: usize, + pub encoder_kwargs: Vec, + pub decoder_kwargs: Vec, + pub quantizer_type: String, + pub quantizer_kwargs: MossAudioTokenizerQuantizerKwargs, + pub reversed_decoder_kwargs: Vec, +} + + +#[derive(Debug, Deserialize)] +pub struct MossAudioTokenizerModuleConfig { + pub module_type: String, + pub patch_size: Option, + pub causal: Option, + pub context_duration: Option, + pub conv_layout: Option, + pub d_model: Option, + pub dim_feedforward: Option, + pub gating: Option, + pub input_dimension: Option, + pub layer_scale: Option, + pub max_period: Option, + pub norm: Option, + pub num_heads: Option, + pub num_layers: Option, + pub output_dimension: Option, + pub positional_embedding: Option, +} + +#[derive(Debug, Deserialize)] +pub struct MossAudioTokenizerQuantizerKwargs { + pub codebook_dim: usize, + pub codebook_loss_weight: f64, + pub codebook_size: usize, + pub commitment_loss_weight: f64, + pub input_dim: usize, + pub num_quantizers: usize, + pub output_dim: usize, + pub quantizer_dropout: f64, + pub quantizer_type: String, + pub rvq_dim: usize, +} + +#[derive(Debug, Deserialize)] +pub struct MossTTSConfig { + pub add_cross_attention: bool, + // Audio Tokenizer Specifics + pub audio_assistant_slot_token_id: u32, + pub audio_codebook_sizes: Vec, + pub audio_end_token_id: u32, + pub audio_pad_token_id: u32, + pub audio_start_token_id: u32, + pub audio_tokenizer_sample_rate: usize, + pub audio_user_slot_token_id: u32, + pub audio_vocab_size: usize, + + // Generation/Model Params (Simplified nullables to Options or defaults if not critical) + pub bad_words_ids: Option>, + pub begin_suppress_tokens: Option>, + pub bos_token_id: Option, + pub chunk_size_feed_forward: usize, + pub cross_attention_hidden_size: Option, + pub decoder_start_token_id: Option, + pub diversity_penalty: f64, + pub do_sample: bool, + pub dtype: String, + pub early_stopping: bool, + pub encoder_no_repeat_ngram_size: usize, + pub eos_token_id: Option, + pub exponential_decay_length_penalty: Option, + pub finetuning_task: Option, + pub forced_bos_token_id: Option, + pub forced_eos_token_id: Option, + + // GPT2 Backbone Config + pub gpt2_config: GPT2Config, + + pub hidden_size: usize, + pub id2label: std::collections::HashMap, + + pub im_end_token_id: u32, + pub im_start_token_id: u32, + pub initializer_range: f64, + pub is_decoder: bool, + pub is_encoder_decoder: bool, + pub label2id: std::collections::HashMap, + + pub length_penalty: f64, + pub local_transformer_attn_implementation: String, + pub local_transformer_layers: usize, + + pub max_length: usize, + pub max_position_embeddings: usize, + pub min_length: usize, + + pub model_architecture: String, + pub model_type: String, + + pub n_vq: usize, + pub no_repeat_ngram_size: usize, + + pub num_beam_groups: usize, + pub num_beams: usize, + pub num_return_sequences: usize, + + pub output_attentions: bool, + pub output_hidden_states: bool, + pub output_scores: bool, + + pub pad_token_id: u32, + pub prefix: Option, + pub problem_type: Option, + // pub pruned_heads: std::collections::HashMap>, + + pub remove_invalid_values: bool, + pub repetition_penalty: f64, + + pub return_dict: bool, + pub return_dict_in_generate: bool, + + pub sep_token_id: Option, + pub suppress_tokens: Option>, + + pub task_specific_params: Option, + + pub temperature: f32, + pub tf_legacy_loss: bool, + pub tie_encoder_decoder: bool, + pub tie_word_embeddings: bool, + + pub tokenizer_class: String, + pub tokenizer_use_fast: bool, + + pub top_k: usize, + pub top_p: f32, + pub torchscript: bool, + + pub typical_p: f64, + + pub use_bfloat16: bool, + pub vocab_size: usize, + +} + + diff --git a/src/models/moss/generate.rs b/src/models/moss/generate.rs new file mode 100644 index 0000000..034774c --- /dev/null +++ b/src/models/moss/generate.rs @@ -0,0 +1,95 @@ +use std::collections::HashMap; + +use crate::{ + models::moss::{ + audio_tokenizer_nano::MossAudioTokenizer, + config::{MossAudioTokenizerConfig, MossTTSConfig}, + processor::MossTTSProcessor, + tts_nano::{MossTTSMode, MossTTSModel}, + }, + utils::{find_type_files, get_device, get_dtype}, +}; +use anyhow::{Result, anyhow}; +use candle_core::{DType, Device, pickle::read_all_with_key}; +use candle_nn::VarBuilder; +use sentencepiece::SentencePieceProcessor; + +pub struct MossTTSGenerate { + pub audio_tokenizer: MossAudioTokenizer, + pub text_tokenizer: SentencePieceProcessor, + pub processor: MossTTSProcessor, + pub model: MossTTSModel, + pub device: Device, +} + +impl MossTTSGenerate { + pub fn init( + tts_path: &str, + audio_tokenizer_path: &str, + device: Option<&Device>, + dtype: Option, + ) -> Result { + let audio_tokenizer_config_path = audio_tokenizer_path.to_string() + "/config.json"; + let audio_tokenizer_cfg: MossAudioTokenizerConfig = + serde_json::from_slice(&std::fs::read(audio_tokenizer_config_path)?)?; + let model_list = find_type_files(audio_tokenizer_path, "safetensors")?; + let audio_dtype = get_dtype(dtype.clone(), &audio_tokenizer_cfg.dtype); + let device = get_device(device); + let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, audio_dtype, &device)? }; + let audio_tokenizer = MossAudioTokenizer::new(vb, &audio_tokenizer_cfg)?; + let text_tokenizer_path = tts_path.to_string() + "/tokenizer.model"; + let text_tokenizer = SentencePieceProcessor::open(text_tokenizer_path) + .map_err(|e| anyhow!(format!("load bpe.model file error:{}", e)))?; + let tts_cfg_path = tts_path.to_string() + "/config.json"; + let tts_cfg: MossTTSConfig = serde_json::from_slice(&std::fs::read(tts_cfg_path)?)?; + let processor = MossTTSProcessor::new( + &tts_cfg, + audio_tokenizer_cfg.sample_rate, + audio_tokenizer_cfg.number_channels, + &text_tokenizer, + )?; + let model_list = find_type_files(tts_path, "bin")?; + let mut dict_to_hashmap = HashMap::new(); + // let cfg_dtype = tts_cfg.dtype.as_str(); + let m_dtype = get_dtype(dtype, "bfloat16"); + for m in model_list { + let dict = read_all_with_key(m, None)?; + for (k, v) in dict { + // println!("key: {}, tensor shape: {:?}", k, v); + dict_to_hashmap.insert(k, v); + } + } + let vb = VarBuilder::from_tensors(dict_to_hashmap, m_dtype, &device); + let model = MossTTSModel::new(vb, &tts_cfg)?; + + Ok(Self { + audio_tokenizer, + text_tokenizer, + processor, + model, + device, + }) + } + + pub fn generate( + &mut self, + text: &str, + prompt_audio_path: Option<&str>, + prompt_text: Option<&str>, + mode: Option, + ) -> Result<()> { + let (mut input_ids, mask) = self.processor.build_inference_input_ids( + text, + prompt_audio_path, + prompt_text, + mode, + &self.audio_tokenizer, + &self.text_tokenizer, + &self.device, + )?; + let _ = self.model.generate(&input_ids, Some(&mask))?; + // println!("input_ids: {}", input_ids); + // println!("mask: {}", mask); + Ok(()) + } +} diff --git a/src/models/moss/mod.rs b/src/models/moss/mod.rs new file mode 100644 index 0000000..bfe491d --- /dev/null +++ b/src/models/moss/mod.rs @@ -0,0 +1,5 @@ +pub mod audio_tokenizer_nano; +pub mod config; +pub mod generate; +pub mod processor; +pub mod tts_nano; diff --git a/src/models/moss/processor.rs b/src/models/moss/processor.rs new file mode 100644 index 0000000..2d87471 --- /dev/null +++ b/src/models/moss/processor.rs @@ -0,0 +1,218 @@ +use crate::{ + models::moss::{ + audio_tokenizer_nano::MossAudioTokenizer, config::MossTTSConfig, tts_nano::MossTTSMode, + }, + tokenizer::sentencepiece_encode_vec, + utils::{audio_utils::load_audio_with_resample, prepare_tts_text}, +}; +use anyhow::{Result, anyhow}; +use candle_core::{Device, Tensor}; +use sentencepiece::SentencePieceProcessor; + +pub struct MossTTSProcessor { + target_sample_rate: usize, + target_channels: usize, + audio_start_token_id: u32, + audio_end_token_id: u32, + audio_user_slot_token_id: u32, + audio_assistant_slot_token_id: u32, + audio_pad_token_id: u32, + n_vq: usize, + prompt_token_ids: Vec, + user_after_ids: Vec, + assistant_ids: Vec, + none_ids: Vec, +} + +impl MossTTSProcessor { + pub fn new( + tts_cfg: &MossTTSConfig, + target_sample_rate: usize, + target_channels: usize, + text_tokenizer: &SentencePieceProcessor, + ) -> Result { + let mut prompt_token_ids = vec![tts_cfg.im_start_token_id]; + let user_role_ids = sentencepiece_encode_vec("user\n", text_tokenizer)?; + prompt_token_ids.extend_from_slice(&user_role_ids); + let user_template_pre_ids = + sentencepiece_encode_vec("\n- Reference(s):\n", text_tokenizer)?; + prompt_token_ids.extend_from_slice(&user_template_pre_ids); + + let user_after_ids = sentencepiece_encode_vec( + "\n- Instruction:\nNone\n- Tokens:\nNone\n- Quality:\nNone\n- Sound Event:\nNone\n- Ambient Sound:\nNone\n- Language:\nNone\n- Text:\n", + text_tokenizer, + )?; + let mut assistant_ids = vec![]; + let user_suffix = sentencepiece_encode_vec("\n", text_tokenizer)?; + assistant_ids.extend_from_slice(&user_suffix); + assistant_ids.push(tts_cfg.im_end_token_id); + let assistant_turn_ids = sentencepiece_encode_vec("\n", text_tokenizer)?; + assistant_ids.extend_from_slice(&assistant_turn_ids); + assistant_ids.push(tts_cfg.im_start_token_id); + let assistant_role_ids = sentencepiece_encode_vec("assistant\n", text_tokenizer)?; + assistant_ids.extend_from_slice(&assistant_role_ids); + + let none_ids = sentencepiece_encode_vec("None", text_tokenizer)?; + + Ok(Self { + target_sample_rate, + target_channels, + audio_start_token_id: tts_cfg.audio_start_token_id, + audio_end_token_id: tts_cfg.audio_end_token_id, + audio_user_slot_token_id: tts_cfg.audio_user_slot_token_id, + audio_assistant_slot_token_id: tts_cfg.audio_assistant_slot_token_id, + audio_pad_token_id: tts_cfg.audio_pad_token_id, + n_vq: tts_cfg.n_vq, + prompt_token_ids, + user_after_ids, + assistant_ids, + none_ids, + }) + } + + fn resolved_mode( + &self, + mode: Option, + has_prompt_text: bool, + has_prompt_audio: bool, + ) -> Result { + let normalized_mode = mode.unwrap_or(MossTTSMode::VoiceClone); + if normalized_mode == MossTTSMode::VoiceClone { + if !has_prompt_audio { + return Err(anyhow!("voice_clone mode requires prompt_audio_path")); + } + if has_prompt_text { + println!("voice_clone mode does not accept prompt_text"); + } + } else { + if has_prompt_text != has_prompt_audio { + return Err(anyhow!( + "continuation mode accepts either target text only, or prompt_text and prompt_audio_path together." + )); + } + } + Ok(normalized_mode) + } + + pub fn build_inference_input_ids( + &self, + text: &str, + prompt_audio_path: Option<&str>, + prompt_text: Option<&str>, + mode: Option, + audio_tokenizer: &MossAudioTokenizer, + text_tokenizer: &SentencePieceProcessor, + device: &Device, + ) -> Result<(Tensor, Tensor)> { + let mode = self.resolved_mode(mode, prompt_text.is_some(), prompt_audio_path.is_some())?; + let audio_code = if let Some(audio_path) = prompt_audio_path { + let audio = load_audio_with_resample( + audio_path, + device, + Some(self.target_sample_rate), + Some(self.target_channels), + )?; + Some(audio_tokenizer.encode_one(&audio)?) + } else { + None + }; + let text = &prepare_tts_text(text)?; + let prompt_text = if let Some(prompt_text) = prompt_text { + Some(prepare_tts_text(prompt_text)?) + } else { + None + }; + // TODO: 长文本段切分 + if mode == MossTTSMode::VoiceClone + && let Some(prompt_audio_codes) = &audio_code + { + let mut prompt_token_ids = vec![]; + prompt_token_ids.extend_from_slice(&self.prompt_token_ids); + prompt_token_ids.push(self.audio_start_token_id); + let prompt_ids_tensor = Self::build_text_raw( + &prompt_token_ids, + self.audio_pad_token_id, + self.n_vq, + device, + )?; + let text_token_ids = sentencepiece_encode_vec(text, text_tokenizer)?; + let mut suffix_token_ids = vec![self.audio_end_token_id]; + suffix_token_ids.extend_from_slice(&self.user_after_ids); + suffix_token_ids.extend_from_slice(&text_token_ids); + suffix_token_ids.extend_from_slice(&self.assistant_ids); + suffix_token_ids.push(self.audio_start_token_id); + let audio_prefix_rows = Self::build_audio_prefix_rows( + &prompt_audio_codes, + self.audio_user_slot_token_id, + device, + )?; + let suffix_rows = Self::build_text_raw( + &suffix_token_ids, + self.audio_pad_token_id, + self.n_vq, + device, + )?; + let input_ids = + Tensor::cat(&[&prompt_ids_tensor, &audio_prefix_rows, &suffix_rows], 0)? + .unsqueeze(0)?; + let (bs, len, _) = input_ids.dims3()?; + let mask = Tensor::ones((bs, len), candle_core::DType::U32, device)?; + Ok((input_ids, mask)) + } else { + let text = if let Some(prompt_text) = prompt_text { + prompt_text + text + } else { + text.to_string() + }; + let text_token_ids = sentencepiece_encode_vec(&text, text_tokenizer)?; + let mut prompt_ids = vec![]; + prompt_ids.extend_from_slice(&self.prompt_token_ids); + prompt_ids.extend_from_slice(&self.none_ids); + prompt_ids.extend_from_slice(&self.user_after_ids); + prompt_ids.extend_from_slice(&text_token_ids); + prompt_ids.extend_from_slice(&self.assistant_ids); + prompt_ids.push(self.audio_start_token_id); + let mut input_ids = + Self::build_text_raw(&prompt_ids, self.audio_pad_token_id, self.n_vq, device)?; + if let Some(prompt_audio_codes) = &audio_code { + let audio_prefix_rows = Self::build_audio_prefix_rows( + &prompt_audio_codes, + self.audio_assistant_slot_token_id, + device, + )?; + input_ids = Tensor::cat(&[&input_ids, &audio_prefix_rows], 0)?; + } + input_ids = input_ids.unsqueeze(0)?; + let (bs, len, _) = input_ids.dims3()?; + let mask = Tensor::ones((bs, len), candle_core::DType::U32, device)?; + Ok((input_ids, mask)) + } + } + + + fn build_audio_prefix_rows( + prompt_audio_codes: &Tensor, + slot_token_id: u32, + device: &Device, + ) -> Result { + let audio_len = prompt_audio_codes.dim(0)?; + let pad_tensor = Tensor::new(slot_token_id, device)?.broadcast_as((audio_len, 1))?; + let rows = Tensor::cat(&[&pad_tensor, prompt_audio_codes], 1)?; + Ok(rows) + } + + fn build_text_raw( + token_ids: &Vec, + audio_pad_token_id: u32, + n_vq: usize, + device: &Device, + ) -> Result { + let id_len = token_ids.len(); + //(1, len) -> (len, 1) + let text_tensor = Tensor::from_slice(token_ids, (id_len, 1), device)?; + // (len, n_vq) + let pad_tensor = Tensor::new(audio_pad_token_id, device)?.broadcast_as((id_len, n_vq))?; + let rows = Tensor::cat(&[text_tensor, pad_tensor], 1)?; + Ok(rows) + } +} diff --git a/src/models/moss/tts_nano.rs b/src/models/moss/tts_nano.rs new file mode 100644 index 0000000..348696b --- /dev/null +++ b/src/models/moss/tts_nano.rs @@ -0,0 +1,256 @@ +use crate::models::{common::sample::simple_sample, gpt2::GPT2Model, moss::config::MossTTSConfig}; +use anyhow::{Result, anyhow}; +use candle_core::{D, IndexOp, Tensor}; +use candle_nn::{Embedding, Linear, Module, VarBuilder, embedding, linear_no_bias}; + +#[derive(PartialEq, Debug)] +pub enum MossTTSMode { + Continuation, + VoiceClone, +} + +pub struct MossTTSModel { + transformer: GPT2Model, + audio_embeddings: Vec, + text_lm_head: Linear, + audio_lm_heads: Vec, + local_transformer: GPT2Model, + audio_assistant_slot_token_id: usize, + audio_end_token_id: usize, + n_vq: usize, + audio_pad_token_id_tensor: Tensor, + audio_codebook_sizes: Vec, + audio_temperature: f64, + audio_top_k: usize, + audio_top_p: f32, + audio_repetition_penalty: f32, +} + +impl MossTTSModel { + pub fn new(vb: VarBuilder, cfg: &MossTTSConfig) -> Result { + let transformer = GPT2Model::new( + vb.pp("transformer"), + cfg.gpt2_config.n_embd, + cfg.gpt2_config.n_head, + cfg.gpt2_config.n_layer, + cfg.gpt2_config.vocab_size, + // cfg.gpt2_config.n_positions, + )?; + let mut audio_embeddings = vec![]; + let audio_embed_vb = vb.pp("audio_embeddings"); + for i in 0..cfg.n_vq { + let embed = embedding( + cfg.audio_codebook_sizes[i], + cfg.gpt2_config.n_embd, + audio_embed_vb.pp(i), + )?; + audio_embeddings.push(embed); + } + let text_lm_head = linear_no_bias( + cfg.gpt2_config.n_embd, + cfg.gpt2_config.vocab_size, + vb.pp("text_lm_head"), + )?; + + let mut audio_lm_heads = vec![]; + let audio_lm_vb = vb.pp("audio_lm_heads"); + for i in 0..cfg.n_vq { + let layer = linear_no_bias( + cfg.gpt2_config.n_embd, + cfg.audio_codebook_sizes[i], + audio_lm_vb.pp(i), + )?; + audio_lm_heads.push(layer); + } + + let mut local_gpt2_cfg = cfg.gpt2_config.clone(); + local_gpt2_cfg.n_layer = cfg.local_transformer_layers; + local_gpt2_cfg.n_positions = cfg.n_vq + 1; + local_gpt2_cfg.n_ctx = cfg.n_vq + 1; + let local_transformer = GPT2Model::new_without_wte( + vb.pp("local_transformer"), + local_gpt2_cfg.n_embd, + local_gpt2_cfg.n_head, + local_gpt2_cfg.n_layer, + local_gpt2_cfg.vocab_size, + // local_gpt2_cfg.n_positions, + )?; + let audio_pad_token_id_tensor = Tensor::new(cfg.audio_pad_token_id, vb.device())?; + // let audio_processor = get_logit_processor(Some(0.8), Some(0.95), Some(25), 34562); + Ok(Self { + transformer, + audio_embeddings, + text_lm_head, + audio_lm_heads, + local_transformer, + audio_assistant_slot_token_id: cfg.audio_assistant_slot_token_id as usize, + audio_end_token_id: cfg.audio_end_token_id as usize, + n_vq: cfg.n_vq, + audio_pad_token_id_tensor, + audio_codebook_sizes: cfg.audio_codebook_sizes.clone(), + audio_temperature: 0.8, + audio_top_k: 25, + audio_top_p: 0.95, + audio_repetition_penalty: 1.2, + }) + } + + fn build_inputs_embeds(&self, input_ids: &Tensor) -> Result { + let text_ids = input_ids.narrow(D::Minus1, 0, 1)?.squeeze(D::Minus1)?; + let mut inputs_embeds = if let Some(wte) = &self.transformer.wte { + wte.forward(&text_ids)? + } else { + return Err(anyhow!("MossTTS transformer wte can not be none")); + }; + for (channel_index, embedding) in self.audio_embeddings.iter().enumerate() { + let channel_ids = input_ids + .narrow(D::Minus1, channel_index + 1, 1)? + .squeeze(D::Minus1)?; + let valid_mask = channel_ids.ne(&self + .audio_pad_token_id_tensor + .broadcast_as(channel_ids.shape())?)?; + let invalid_mask = channel_ids.lt(&channel_ids.zeros_like()?)?; + let embedding_nums = Tensor::new( + self.audio_codebook_sizes[channel_index] as u32, + input_ids.device(), + )?; + let invalid_mask1 = + channel_ids.ge(&embedding_nums.broadcast_as(channel_ids.shape())?)?; + let invalid_mask = valid_mask + .minimum(&invalid_mask.maximum(&invalid_mask1)?)? + .to_dtype(candle_core::DType::U32)?; + if invalid_mask.sum_all()?.to_scalar::()? > 0 { + return Err(anyhow!("Found out-of-range audio token ids for channel")); + } + let safe_ids = valid_mask.where_cond(&channel_ids, &channel_ids.zeros_like()?)?; + let audio_embeds = embedding.forward(&safe_ids)?; + let audio_embeds = audio_embeds.broadcast_mul( + &valid_mask + .unsqueeze(D::Minus1)? + .to_dtype(audio_embeds.dtype())?, + )?; + inputs_embeds = inputs_embeds.add(&audio_embeds)?; + } + Ok(inputs_embeds) + } + + fn sample_next_assistant_text_token(&self, logits: &Tensor) -> Result { + let logits = logits.squeeze(0)?.squeeze(0)?; + let slot_token_id_logit = logits + .i(self.audio_assistant_slot_token_id)? + .to_dtype(candle_core::DType::F32)? + .to_scalar::()?; + let end_token_id_logit = logits + .i(self.audio_end_token_id)? + .to_dtype(candle_core::DType::F32)? + .to_scalar::()?; + println!( + "slot_token_id: {} logit: {slot_token_id_logit}", + self.audio_assistant_slot_token_id + ); + println!( + "end_token_id: {} logit: {end_token_id_logit}", + self.audio_end_token_id + ); + if slot_token_id_logit > end_token_id_logit { + Ok(self.audio_assistant_slot_token_id) + } else { + Ok(self.audio_end_token_id) + } + } + + fn build_generation_row(&self, audio_token_ids: &Tensor) -> Result { + let slot = Tensor::from_slice( + &[self.audio_assistant_slot_token_id as u32], + (1, 1, 1), + audio_token_ids.device(), + )?; + let audio_token_ids = audio_token_ids.unsqueeze(0)?.unsqueeze(0)?; + Ok(Tensor::cat(&[&slot, &audio_token_ids], D::Minus1)?) + } + + pub fn generate(&mut self, input_ids: &Tensor, mask: Option<&Tensor>) -> Result<()> { + let sample_len = 2; + let mut seqlen_offset = 0; + let mut seq_len = input_ids.dim(1)?; + let mut generated_frames = vec![]; + let mut current_model_input_ids = input_ids.clone(); + for step_index in 0..sample_len { + // println!("current_model_input_ids: {:?}", current_model_input_ids); + let inputs_embeds = self.build_inputs_embeds(¤t_model_input_ids)?; + let outputs = self.transformer.forward(&inputs_embeds, seqlen_offset)?; + // println!("transformer-----------------------"); + let outputs_len = outputs.dim(1)?; + let global_hidden_state = outputs.narrow(1, outputs_len - 1, 1)?; + // println!("global_hidden_state: {}", global_hidden_state); + let mut local_positions = 0usize; + let local_outputs = self + .local_transformer + .forward(&global_hidden_state, local_positions)?; + // println!("local_outputs-----------------------"); + let local_len = local_outputs.dim(1)?; + let local_hidden_states = local_outputs.narrow(1, local_len - 1, 1)?; + // println!("local_hidden_states: {}", local_hidden_states); + let text_logits = self.text_lm_head.forward(&local_hidden_states)?; + // println!("text_logits: {}", text_logits.i((0, 0, 0..100))?); + println!("step_index: {}", step_index); + let next_text_token = self.sample_next_assistant_text_token(&text_logits)?; + if next_text_token == self.audio_end_token_id { + self.local_transformer.clear_kv_cache(); + break; + } + let mut next_frame_tokens = vec![]; + let mut current_local_input = if let Some(wte) = &self.transformer.wte { + wte.forward(&Tensor::from_slice( + &[next_text_token as u32], + (1, 1), + input_ids.device(), + )?)? + } else { + return Err(anyhow!("MossTTS GPT2 wte can not be none")); + }; + for channel_index in 0..self.n_vq { + local_positions += 1; + let local_outputs = self + .local_transformer + .forward(¤t_local_input, local_positions)?; + let local_len = local_outputs.dim(1)?; + let local_hidden_states = local_outputs.narrow(1, local_len - 1, 1)?; + // println!("local_hidden_states: {local_hidden_states}"); + let channel_logits = (&self.audio_lm_heads[channel_index]) + .forward(&local_hidden_states)? + .squeeze(0)? + .squeeze(0)?; + // println!("channel_logits: {}", channel_logits.i(0..100)?); + let arg_max = channel_logits.argmax(0)?; + println!("arg_max: {}", arg_max); + let channel_token = simple_sample( + &channel_logits, + true, + Some(self.audio_temperature), + Some(self.audio_top_k), + Some(self.audio_top_p), + Some(&next_frame_tokens), + self.audio_repetition_penalty, + None, + )?; + println!("channel_token: {channel_token}"); + next_frame_tokens.push(channel_token); + current_local_input = (&self.audio_embeddings[channel_index]).forward( + &Tensor::from_slice(&[channel_token], (1, 1), input_ids.device())?, + )?; + // println!("current_local_input: {current_local_input}"); + } + self.local_transformer.clear_kv_cache(); + let next_frame = Tensor::new(next_frame_tokens, input_ids.device())?; + // println!("next_frame: {next_frame}"); + current_model_input_ids = self.build_generation_row(&next_frame)?; + seqlen_offset += seq_len; + seq_len = 1; + generated_frames.push(next_frame); + } + let audio_token_ids = Tensor::stack(&generated_frames, 0)?; + println!("audio_token_ids: {audio_token_ids}"); + Ok(()) + } +} diff --git a/src/models/qwen2_5vl/generate.rs b/src/models/qwen2_5vl/generate.rs index 3203c09..6ddaf1e 100644 --- a/src/models/qwen2_5vl/generate.rs +++ b/src/models/qwen2_5vl/generate.rs @@ -1,6 +1,6 @@ use std::time::Instant; -use crate::models::common::generate::get_logit_processor; +use crate::models::common::sample::get_logit_processor; use crate::params::chat::{ ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse, }; diff --git a/src/models/qwen3_asr/generate.rs b/src/models/qwen3_asr/generate.rs index 5af3bd4..65bdaab 100644 --- a/src/models/qwen3_asr/generate.rs +++ b/src/models/qwen3_asr/generate.rs @@ -3,8 +3,8 @@ use std::time::Instant; use crate::{ models::common::{ MultiModalData, - generate::{GenerationContext, generate_generic_text, get_logit_processor}, - modules::{AsrResult, VadFrameResult}, + generate::{GenerationContext, generate_generic_text}, + modules::{AsrResult, VadFrameResult}, sample::get_logit_processor, }, params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}, utils::response_utils::{build_chunk_response_with_usage, build_completion_response_with_time}, diff --git a/src/models/qwen3_asr/processor.rs b/src/models/qwen3_asr/processor.rs index 859022c..2570b62 100644 --- a/src/models/qwen3_asr/processor.rs +++ b/src/models/qwen3_asr/processor.rs @@ -82,7 +82,7 @@ impl Qwen3AsrProcessor { } pub fn extract_audio_vec(&self, mes: &ChatCompletionParameters) -> Result> { - let audio_tensors = extract_audios(mes, &self.device, Some(self.sample_rate))?; + let audio_tensors = extract_audios(mes, &self.device, Some(self.sample_rate), Some(1))?; audio_tensors.iter().map(float_range_normalize).collect() } diff --git a/src/models/voxcpm/model.rs b/src/models/voxcpm/model.rs index 16265b0..c2c8644 100644 --- a/src/models/voxcpm/model.rs +++ b/src/models/voxcpm/model.rs @@ -534,7 +534,8 @@ impl VoxCPMModel { let audio_start = Tensor::new(vec![self.audio_start_token], &self.device)?; let text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?; let text_length = text_token.dim(0)?; - let mut audio = load_audio_with_resample(&path, &self.device, Some(self.sample_rate))?; + let mut audio = + load_audio_with_resample(&path, &self.device, Some(self.sample_rate), Some(1))?; let patch_len = self.patch_size * self.chunk_size; if audio.dim(1)? % patch_len != 0 { audio = @@ -574,7 +575,8 @@ impl VoxCPMModel { let audio_start = Tensor::new(vec![self.audio_start_token], &self.device)?; let text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?; let text_length = text_token.dim(0)?; - let mut audio = load_audio_with_resample(&path, &self.device, Some(self.sample_rate))?; + let mut audio = + load_audio_with_resample(&path, &self.device, Some(self.sample_rate), Some(1))?; let patch_len = self.patch_size * self.chunk_size; if audio.dim(1)? % patch_len != 0 { audio = @@ -841,8 +843,12 @@ impl VoxCPMModel { ) -> Result> { let text_token = self.tokenizer.encode(prompt_text)?; let text_token = Tensor::from_slice(&text_token, text_token.len(), &self.device)?; - let mut audio = - load_audio_with_resample(&prompt_wav_path, &self.device, Some(self.sample_rate))?; + let mut audio = load_audio_with_resample( + &prompt_wav_path, + &self.device, + Some(self.sample_rate), + Some(1), + )?; let patch_len = self.patch_size * self.chunk_size; if audio.dim(1)? % patch_len != 0 { audio = audio.pad_with_zeros(D::Minus1, 0, patch_len - audio.dim(1)? % patch_len)?; diff --git a/src/models/voxcpm_refact/processor.rs b/src/models/voxcpm_refact/processor.rs index 00bdf60..b53a663 100644 --- a/src/models/voxcpm_refact/processor.rs +++ b/src/models/voxcpm_refact/processor.rs @@ -38,8 +38,12 @@ impl VoxCPMProcessor { audio_vae: &AudioVAE, ) -> Result> { let (text_token, _) = tokenizer.encode_tensor(prompt_text, &self.device)?; - let mut audio = - load_audio_with_resample(&prompt_wav_path, &self.device, Some(self.sample_rate))?; + let mut audio = load_audio_with_resample( + &prompt_wav_path, + &self.device, + Some(self.sample_rate), + Some(1), + )?; let patch_len = self.patch_size * self.chunk_size; if audio.dim(1)? % patch_len != 0 { audio = audio.pad_with_zeros(D::Minus1, 0, patch_len - audio.dim(1)? % patch_len)?; @@ -74,7 +78,8 @@ impl VoxCPMProcessor { let mut text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?; let (audio_feat, audio_mask) = if let Some(path) = prompt_wav_path { - let mut audio = load_audio_with_resample(&path, &self.device, Some(self.sample_rate))?; + let mut audio = + load_audio_with_resample(&path, &self.device, Some(self.sample_rate), Some(1))?; let patch_len = self.patch_size * self.chunk_size; if audio.dim(1)? % patch_len != 0 { audio = diff --git a/src/position_embed/rope.rs b/src/position_embed/rope.rs index 0f18e01..300ec25 100644 --- a/src/position_embed/rope.rs +++ b/src/position_embed/rope.rs @@ -1,4 +1,4 @@ -use anyhow::{Result, anyhow}; +use anyhow::Result; use candle_core::{D, DType, Device, IndexOp, Tensor}; use candle_transformers::models::deepseek2::SplitOp; @@ -21,6 +21,22 @@ pub fn rotate_half(x: &Tensor) -> Result { Ok(rotate_x) } +pub fn rotate_half_interleave(x: &Tensor) -> Result { + let x_rank = x.rank(); + let x_dim = x.dims(); + let half_dim = x_dim[x_rank - 1] / 2; + let mut x_reshape = x_dim[0..x_rank - 1].to_vec(); + x_reshape.push(half_dim); + x_reshape.push(2); + let x = x.reshape(x_reshape)?; + let even = x.narrow(D::Minus1, 0, 1)?; + let odd = x.narrow(D::Minus1, 1, 1)?.affine(-1.0, 0.0)?; + let rotate_x = Tensor::cat(&[&odd, &even], D::Minus1)? + .reshape(x_dim)? + .contiguous()?; + Ok(rotate_x) +} + pub fn apply_multimodel_rotary_pos_emb( q: &Tensor, k: &Tensor, @@ -115,6 +131,44 @@ pub fn apply_rotary_pos_emb( Ok((q_embed, k_embed)) } +pub fn apply_rotary_pos_emb_interleave( + q: &Tensor, + k: &Tensor, + cos: &Tensor, + sin: &Tensor, + tof32: bool, +) -> Result<(Tensor, Tensor)> { + // sin/cos: to (bs, 1, seq_len, head_dim) + // q/k: (bs, n_head, seq_len, head_dim) + let mut cos = cos.clone(); + let mut sin = sin.clone(); + if cos.rank() == 2 { + // (seq_len, head_dim) -> (1, 1, seq_len, head_dim) + cos = cos.unsqueeze(0)?.unsqueeze(0)?; + sin = sin.unsqueeze(0)?.unsqueeze(0)?; + } + if cos.rank() == 3 { + // (bs, seq_len, head_dim) -> (bs, 1, seq_len, head_dim) + cos = cos.unsqueeze(1)?; + sin = sin.unsqueeze(1)?; + } + let orig_dtype = q.dtype(); + let q = if tof32 { &q.to_dtype(DType::F32)? } else { q }; + let k = if tof32 { &k.to_dtype(DType::F32)? } else { k }; + let cos = cos.to_dtype(q.dtype())?; + let sin = sin.to_dtype(q.dtype())?; + + let q_embed = q + .broadcast_mul(&cos)? + .add(&rotate_half_interleave(q)?.broadcast_mul(&sin)?)? + .to_dtype(orig_dtype)?; + let k_embed = k + .broadcast_mul(&cos)? + .add(&rotate_half_interleave(k)?.broadcast_mul(&sin)?)? + .to_dtype(orig_dtype)?; + Ok((q_embed, k_embed)) +} + pub fn glm_asr_apply_rotary_pos_emb( q: &Tensor, k: &Tensor, @@ -258,66 +312,46 @@ pub fn glm_ocr_apply_rotary_pos_emb( Ok((q_embed, k_embed)) } -pub fn roformer_rotate(x: &Tensor) -> Result { - let dims = x.dims(); - let last_dim = dims - .last() - .ok_or(anyhow!("Input tensor must have at least one dimension"))?; - if last_dim % 2 != 0 { - return Err(anyhow!( - "Last dimension size must be even, got {}", - last_dim - )); - } - let new_dims: Vec = dims[..dims.len() - 1] - .iter() - .copied() - .chain([last_dim / 2, 2]) - .collect(); - let x_reshape = x.reshape(new_dims)?; - let x_chunks = x_reshape.chunk(2, D::Minus1)?; - let x1 = &x_chunks[0]; - let x2 = &x_chunks[1]; - // let x1 = x_reshape.narrow(D::Minus1, 0, 1)?; - // let x2 = x_reshape.narrow(D::Minus1, 1, 1)?; - let x2_neg = x2.affine(-1.0, 0.0)?; - let rotate_x = Tensor::cat(&[&x2_neg, x1], D::Minus1)?; - Ok(rotate_x.flatten(D::Minus2, D::Minus1)?) -} - pub fn apply_rotary_pos_emb_roformer( q: &Tensor, k: &Tensor, cos: &Tensor, sin: &Tensor, - tof32: bool, ) -> Result<(Tensor, Tensor)> { - let mut cos = cos.clone(); - let mut sin = sin.clone(); - if cos.rank() == 2 { - // (seq_len, head_dim) -> (1, 1, seq_len, head_dim) - cos = cos.unsqueeze(0)?.unsqueeze(0)?; - sin = sin.unsqueeze(0)?.unsqueeze(0)?; - } - if cos.rank() == 3 { - // (bs, seq_len, head_dim) -> (bs, 1, seq_len, head_dim) - cos = cos.unsqueeze(1)?; - sin = sin.unsqueeze(1)?; - } - let orig_dtype = q.dtype(); - let q = if tof32 { &q.to_dtype(DType::F32)? } else { q }; - let k = if tof32 { &k.to_dtype(DType::F32)? } else { k }; - let cos = cos.to_dtype(q.dtype())?; - let sin = sin.to_dtype(q.dtype())?; - let q_embed = q - .broadcast_mul(&cos)? - .add(&roformer_rotate(q)?.broadcast_mul(&sin)?)? - .to_dtype(orig_dtype)?; - let k_embed = k - .broadcast_mul(&cos)? - .add(&roformer_rotate(k)?.broadcast_mul(&sin)?)? - .to_dtype(orig_dtype)?; - Ok((q_embed, k_embed)) + let ori_dtype = q.dtype(); + let (bs, n_head, seq_len, dim) = q.dims4()?; + let half_dim = dim / 2; + let rotr = cos + .narrow(D::Minus1, 0, half_dim)? + .to_dtype(candle_core::DType::F32)?; + let roti = sin + .narrow(D::Minus1, 0, half_dim)? + .to_dtype(candle_core::DType::F32)?; + let q = q + .reshape((bs, n_head, seq_len, half_dim, 2))? + .to_dtype(candle_core::DType::F32)?; + let qr = q.narrow(D::Minus1, 0, 1)?.squeeze(D::Minus1)?; + let qi = q.narrow(D::Minus1, 1, 1)?.squeeze(D::Minus1)?; + + let k = k + .reshape((bs, n_head, seq_len, half_dim, 2))? + .to_dtype(candle_core::DType::F32)?; + let kr = k.narrow(D::Minus1, 0, 1)?.squeeze(D::Minus1)?; + let ki = k.narrow(D::Minus1, 1, 1)?.squeeze(D::Minus1)?; + + let qor = qr.broadcast_mul(&rotr)?.sub(&qi.broadcast_mul(&roti)?)?; + let qoi = qr.broadcast_mul(&roti)?.add(&qi.broadcast_mul(&rotr)?)?; + + let kor = kr.broadcast_mul(&rotr)?.sub(&ki.broadcast_mul(&roti)?)?; + let koi = kr.broadcast_mul(&roti)?.add(&ki.broadcast_mul(&rotr)?)?; + + let q = Tensor::stack(&[qor, qoi], D::Minus1)? + .reshape((bs, n_head, seq_len, dim))? + .to_dtype(ori_dtype)?; + let k = Tensor::stack(&[kor, koi], D::Minus1)? + .reshape((bs, n_head, seq_len, dim))? + .to_dtype(ori_dtype)?; + Ok((q, k)) } #[derive(Debug, Clone)] @@ -554,7 +588,6 @@ impl RoPE { pub fn new(dim: usize, theta_base: f32, device: &Device) -> Result { let inv_freq = compute_default_rope_parameters(dim, theta_base); let inv_freq = Tensor::from_slice(&inv_freq, (1, inv_freq.len()), device)?; - Ok(Self { inv_freq }) } pub fn forward( @@ -577,6 +610,35 @@ impl RoPE { let sin = emb.sin()?; Ok((cos, sin)) } + pub fn forward_repeat_interleave( + &self, + seqlen_offset: usize, + seq_len: usize, + device: &Device, + ) -> Result<(Tensor, Tensor)> { + let positions = Tensor::arange( + seqlen_offset as f32, + (seqlen_offset + seq_len) as f32, + self.inv_freq.device(), + )? + .reshape((seq_len, 1))?; // (seq_len, 1) + let freqs = positions.matmul(&self.inv_freq)?; // (seq_len, dim / 2) + let cos = freqs + .cos()? + .unsqueeze(D::Minus1)? + .repeat((1, 1, 2))? + .flatten_from(D::Minus2)? + .contiguous()? + .to_device(device)?; + let sin = freqs + .sin()? + .unsqueeze(D::Minus1)? + .repeat((1, 1, 2))? + .flatten_from(D::Minus2)? + .contiguous()? + .to_device(device)?; + Ok((cos, sin)) + } } pub fn get_xd_cos_sin( diff --git a/src/tokenizer/mod.rs b/src/tokenizer/mod.rs index 7d4c165..c698732 100644 --- a/src/tokenizer/mod.rs +++ b/src/tokenizer/mod.rs @@ -120,15 +120,22 @@ impl TokenizerModel { } } +pub fn sentencepiece_encode_vec( + text: &str, + tokenizer: &SentencePieceProcessor, +) -> Result> { + let tokens = tokenizer + .encode(text) + .map_err(|e| anyhow!(format!("tokenizer encode error:{}", e)))?; + Ok(tokens.iter().map(|p| p.id).collect::>()) +} + pub fn sentencepiece_encode( text: &str, tokenizer: &SentencePieceProcessor, device: &Device, ) -> Result { - let tokens = tokenizer - .encode(text) - .map_err(|e| anyhow!(format!("tokenizer encode error:{}", e)))?; - let token_ids = tokens.iter().map(|p| p.id).collect::>(); + let token_ids = sentencepiece_encode_vec(text, tokenizer)?; let tokens_t = Tensor::new(token_ids, device)?.unsqueeze(0)?; Ok(tokens_t) } diff --git a/src/utils/audio_utils.rs b/src/utils/audio_utils.rs index abd1131..c1d952e 100644 --- a/src/utils/audio_utils.rs +++ b/src/utils/audio_utils.rs @@ -470,7 +470,14 @@ pub fn get_audio_format_from_bytes(bytes: &[u8]) -> Result { } } -pub fn load_audio_use_symphonia(audio_vec: Vec, device: &Device) -> Result<(Tensor, usize)> { +/// return +/// audio shape: (channel, audio_len) +/// sample_rate: usize +pub fn load_audio_use_symphonia( + audio_vec: Vec, + device: &Device, + target_channels: usize, +) -> Result<(Tensor, usize)> { let extension = get_audio_format_from_bytes(&audio_vec)?; let content = Cursor::new(audio_vec); let mss = MediaSourceStream::new(Box::new(content), Default::default()); @@ -558,16 +565,26 @@ pub fn load_audio_use_symphonia(audio_vec: Vec, device: &Device) -> Result<( } } let mut audio_tensor = Tensor::new(all_samples, device)?; - if channels > 1 { - // 对channel通道求平均, channel维度变为1 - audio_tensor = audio_tensor.mean_keepdim(0)?; + if target_channels == channels { + return Ok((audio_tensor, sample_rate as usize)); } - Ok((audio_tensor, sample_rate as usize)) -} -pub fn load_audio(path: &str, device: &Device) -> Result<(Tensor, usize)> { - let audio_vec = get_audio_bytes_vec(path)?; - load_audio_use_symphonia(audio_vec, device) + audio_tensor = if channels == 1 { + // Mono to Multi-channel: Repeat + audio_tensor.repeat((target_channels, 1))? + } else if target_channels == 1 { + // Multi-channel to Mono: Mean + audio_tensor.mean_keepdim(0)? + } else { + // Unsupported conversion (e.g., Stereo to 5.1) + return Err(anyhow!( + "target_channels: {}, audio channels: {}, can't change directly", + target_channels, + channels + )); + }; + + Ok((audio_tensor, sample_rate as usize)) } pub fn resample_audio_from_vec_f32( @@ -599,12 +616,14 @@ pub fn resample_audio_from_vec_f32( Ok(audio) } +/// return shape: (channel, audio_len) pub fn resample_audio_from_bytes( audio_vec: Vec, device: &Device, target_sample_rate: Option, + target_channels: usize, ) -> Result { - let (mut audio, sr) = load_audio_use_symphonia(audio_vec, device)?; + let (mut audio, sr) = load_audio_use_symphonia(audio_vec, device, target_channels)?; if let Some(target_sample_rate) = target_sample_rate && target_sample_rate != sr { @@ -613,16 +632,19 @@ pub fn resample_audio_from_bytes( Ok(audio) } +/// return shape: (channel, audio_len) pub fn load_audio_with_resample( path: &str, device: &Device, target_sample_rate: Option, + target_channels: Option, ) -> Result { // hound 只支持wav文件 // let audio_path = get_audio_path(path)?; // let (mut audio, sr) = load_audio_use_hound(audio_path, device)?; + let target_channels = target_channels.unwrap_or(1); let audio_vec = get_audio_bytes_vec(path)?; - resample_audio_from_bytes(audio_vec, device, target_sample_rate) + resample_audio_from_bytes(audio_vec, device, target_sample_rate, target_channels) } pub fn save_wav(audio: &Tensor, save_path: &str, sample_rate: u32) -> Result<()> { @@ -692,12 +714,13 @@ pub fn extract_audios( mes: &ChatCompletionParameters, device: &Device, target_sample_rate: Option, + target_channels: Option, ) -> Result> { let audio_url_vec = extract_audio_url(mes); // 并行加载音频 audio_url_vec .par_iter() - .map(|url| load_audio_with_resample(url, device, target_sample_rate)) + .map(|url| load_audio_with_resample(url, device, target_sample_rate, target_channels)) .collect() // #[cfg(not(feature = "ffmpeg"))] // { diff --git a/src/utils/mod.rs b/src/utils/mod.rs index d5325bf..6c3a3a4 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -721,6 +721,66 @@ pub fn bucketize(input: usize, boundaries: &[usize]) -> Result { // Ok(index) } +pub fn contains_cjk(text: &str) -> bool { + for ch in text.chars() { + let c = ch as u32; + if (0x4e00..=0x9fff).contains(&c) // CJK Unified Ideographs + || (0x3400..=0x4dbf).contains(&c) // CJK Unified Ideographs Extension A + || (0x3040..=0x30ff).contains(&c) // Hiragana and Katakana + || (0xac00..=0xd7af).contains(&c) // Hangul Syllables + { + return true; + } + } + false +} + +pub fn prepare_tts_text(text: &str) -> Result { + let mut normalized_text = text.trim().to_string(); + if normalized_text.eq("") { + return Err(anyhow!("Text cannot be empty.")) + } + normalized_text = normalized_text.replace('\n', " ").replace('\r', " "); + while normalized_text.contains(" ") { + normalized_text = normalized_text.replace(" ", " "); + } + + if contains_cjk(&normalized_text) { + let cjk_end_punctuations = ['。', '!', '?', '…', '.', '!', '?']; + if !normalized_text.ends_with(|c: char| cjk_end_punctuations.contains(&c)) { + normalized_text.push('。'); + } + return Ok(normalized_text); + } + + // Non-CJK (English/Western) logic + // Capitalize first letter if it's lowercase alphabetic + if let Some(first_char) = normalized_text.chars().next() { + if first_char.is_ascii_lowercase() { + let mut chars = normalized_text.chars(); + chars.next(); // consume first char + let rest: String = chars.collect(); + normalized_text = format!("{}{}", first_char.to_ascii_uppercase(), rest); + } + } + + // Add period if ends with alphanumeric + if let Some(last_char) = normalized_text.chars().last() { + if last_char.is_alphanumeric() { + normalized_text.push('.'); + } + } + + // Add padding if less than 5 words + // Split by whitespace to count words + let word_count = normalized_text.split_whitespace().count(); + if word_count < 5 { + normalized_text = format!(" {}", normalized_text); // 8 spaces + } + + Ok(normalized_text) +} + #[cfg(test)] mod tests { use super::*; diff --git a/tests/config_tests.rs b/tests/config_tests.rs index b13fdd6..5e7b202 100644 --- a/tests/config_tests.rs +++ b/tests/config_tests.rs @@ -1,13 +1,5 @@ use aha::models::{ - deepseek_ocr::config::DeepseekOCRConfig, - hunyuan_ocr::config::HunYuanVLConfig, - lfm2::config::Lfm2Config, - lfm2vl::config::{Lfm2ProcessorConfig, Lfm2VLConfig}, - minicpm4::config::MiniCPM4Config, - paddleocr_vl::config::PaddleOCRVLConfig, - qwen2_5vl::config::Qwen2_5VLConfig, - qwen3vl::config::Qwen3VLConfig, - voxcpm::config::VoxCPMConfig, + deepseek_ocr::config::DeepseekOCRConfig, hunyuan_ocr::config::HunYuanVLConfig, lfm2::config::Lfm2Config, lfm2vl::config::{Lfm2ProcessorConfig, Lfm2VLConfig}, minicpm4::config::MiniCPM4Config, moss::config::{MossAudioTokenizerConfig, MossTTSConfig}, paddleocr_vl::config::PaddleOCRVLConfig, qwen2_5vl::config::Qwen2_5VLConfig, qwen3vl::config::Qwen3VLConfig, voxcpm::config::VoxCPMConfig }; use anyhow::Result; @@ -117,3 +109,23 @@ fn lfm2vl_config() -> Result<()> { println!("{:?}", processor_config); Ok(()) } + +#[test] +fn moss_audio_tokenizer_config() -> Result<()> { + // cargo test -F cuda --test config_tests moss_audio_tokenizer_config -r -- --nocapture + let model_path = "/home/jhq/.aha/openmoss/MOSS-Audio-Tokenizer-Nano/"; + let config_path = model_path.to_string() + "/config.json"; + let config: MossAudioTokenizerConfig = serde_json::from_slice(&std::fs::read(config_path)?)?; + println!("{:?}", config); + Ok(()) +} + +#[test] +fn moss_tts_config() -> Result<()> { + // cargo test -F cuda --test config_tests moss_tts_config -r -- --nocapture + let model_path = "/home/jhq/.aha/openmoss/MOSS-TTS-Nano/"; + let config_path = model_path.to_string() + "/config.json"; + let config: MossTTSConfig = serde_json::from_slice(&std::fs::read(config_path)?)?; + println!("{:?}", config); + Ok(()) +} \ No newline at end of file diff --git a/tests/test_moss_tts.rs b/tests/test_moss_tts.rs new file mode 100644 index 0000000..9964783 --- /dev/null +++ b/tests/test_moss_tts.rs @@ -0,0 +1,20 @@ +use aha::models::moss::generate::MossTTSGenerate; +use anyhow::Result; + +#[test] +fn moss_tts() -> Result<()> { + // RUST_BACKTRACE=1 cargo test -F cuda --test test_moss_tts moss_tts -r -- --nocapture + let save_dir = + aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; + let tts_path = format!("{}/openmoss/MOSS-TTS-Nano/", save_dir); + let audio_tokenizer_path = format!("{}/openmoss/MOSS-Audio-Tokenizer-Nano/", save_dir); + let mut model = MossTTSGenerate::init(&tts_path, &audio_tokenizer_path, None, None)?; + let _ = model.generate( + "您好啊,吃饭了吗,吃的啥啊中午", + Some("file://./assets/audio/jiangjiang.wav"), + Some("哈喽大家好,我是蒋蒋"), + Some(aha::models::moss::tts_nano::MossTTSMode::Continuation), + // None, + )?; + Ok(()) +} diff --git a/tests/weight_test.rs b/tests/weight_test.rs index 90b4938..9e80a3f 100644 --- a/tests/weight_test.rs +++ b/tests/weight_test.rs @@ -425,3 +425,30 @@ fn silero_vad_weight() -> Result<()> { Ok(()) } + +#[test] +fn moss_tts_nano_weight() -> Result<()> { + // cargo test -F cuda --test weight_test moss_tts_nano_weight -r -- --nocapture + let save_dir = + aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; + let model_path = format!("{}/openmoss/MOSS-TTS-Nano/pytorch_model.bin", save_dir); + let dict = read_all_with_key(&model_path, None)?; + for (k, v) in dict { + println!("key: {}, tensor shape: {:?}", k, v); + } + Ok(()) +} + +#[test] +fn moss_audio_tokenizer_nano_weight() -> Result<()> { + // cargo test -F cuda --test weight_test moss_audio_tokenizer_nano_weight -r -- --nocapture + let save_dir = + aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; + let model_path = format!("{}/openmoss/MOSS-Audio-Tokenizer-Nano/model-00001-of-00001.safetensors", save_dir); + let device = get_device(None); + let weights = safetensors::load(model_path, &device)?; + for (key, tensor) in weights.iter() { + println!("=== {} === {:?}", key, tensor); + } + Ok(()) +} \ No newline at end of file