diff --git a/Cargo.lock b/Cargo.lock index 583110f..70850f9 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -885,7 +885,6 @@ dependencies = [ ] [[package]] -<<<<<<< HEAD name = "cpufeatures" version = "0.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -895,23 +894,6 @@ dependencies = [ ] [[package]] -name = "crc" -version = "3.3.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9710d3b3739c2e349eb44fe848ad0b7c8cb1e42bd87ee49371df2f7acaf3e675" -dependencies = [ - "crc-catalog", -] - -[[package]] -name = "crc-catalog" -version = "2.4.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "19d374276b40fb8bbdee95aef7c7fa6b5316ec764510eb64b8dd0e2ed0d7e7f5" - -[[package]] -======= ->>>>>>> main name = "crc32fast" version = "1.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" diff --git a/src/exec/voxcpm.rs b/src/exec/voxcpm.rs index 3673fb0..35ae76d 100644 --- a/src/exec/voxcpm.rs +++ b/src/exec/voxcpm.rs @@ -51,7 +51,7 @@ impl ExecModel for VoxCPMExec { }; let sample_rate = voxcpm_generate.sample_rate(); - crate::utils::audio_utils::save_wav(&audio, &output_path, sample_rate as u32)?; + crate::utils::audio_utils::save_wav_mono(&audio, &output_path, sample_rate as u32)?; println!("Output saved to: {}", output_path); diff --git a/src/models/common/generate.rs b/src/models/common/generate.rs index f9383d3..7e1c1ed 100644 --- a/src/models/common/generate.rs +++ b/src/models/common/generate.rs @@ -1,12 +1,15 @@ use anyhow::Result; use candle_core::{DType, Device, Tensor}; -use candle_transformers::generation::{LogitsProcessor}; +use candle_transformers::generation::LogitsProcessor; use rocket::async_stream::stream; use rocket::futures::Stream; use std::time::Instant; use crate::{ - models::common::{InferenceModel, MultiModalData, sample::{use_repeat_penalty, get_logit_processor}}, + models::common::{ + InferenceModel, MultiModalData, + sample::{get_logit_processor, use_repeat_penalty}, + }, params::chat::{ChatCompletionChunkResponse, ChatCompletionResponse}, tokenizer::TokenizerModel, utils::response_utils::{ diff --git a/src/models/common/modules.rs b/src/models/common/modules.rs index b0f877e..61157ae 100644 --- a/src/models/common/modules.rs +++ b/src/models/common/modules.rs @@ -1393,3 +1393,16 @@ pub fn quick_gelu(xs: &Tensor) -> Result { let x = sigmoid(&x)?; Ok(xs.mul(&x)?) } + +pub fn new_gelu(xs: &Tensor) -> Result { + let sqrt_two_over_pi = (2.0f64 / std::f64::consts::PI).sqrt(); + Ok(xs + .powf(3.0)? + .affine(0.044715, 0.0)? + .add(xs)? + .affine(sqrt_two_over_pi, 0.0)? + .tanh()? + .affine(1.0, 1.0)? + .mul(xs)? + .affine(0.5, 0.0)?) +} diff --git a/src/models/common/sample.rs b/src/models/common/sample.rs index 7f117ac..76917f4 100644 --- a/src/models/common/sample.rs +++ b/src/models/common/sample.rs @@ -2,7 +2,7 @@ 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}; +use rand::distr::Distribution; pub fn get_logit_processor( temperature: Option, @@ -43,7 +43,7 @@ pub fn use_repeat_penalty( logits: &Tensor, context: &[u32], ) -> Result { - if repeat_penalty == 1.0 || repeat_last_n.map_or(false, |n| n == 0) { + if repeat_penalty == 1.0 || repeat_last_n == Some(0) { Ok(logits.clone()) } else { let start_at = if let Some(last_n) = repeat_last_n { @@ -52,13 +52,24 @@ pub fn use_repeat_penalty( 0 }; Ok(candle_transformers::utils::apply_repeat_penalty( - &logits, + logits, repeat_penalty, &context[start_at..], )?) } } +pub fn sample_weighted(prs: &[f32]) -> Result { + let mut rng = rand::rng(); + let dist = rand::distr::weighted::WeightedIndex::new(prs).map_err(|e| { + anyhow!(format!( + "simple_sampel new rand::distr::weighted::WeightedIndex Failed: {}", + e + )) + })?; + Ok(dist.sample(&mut rng) as u32) +} + /// logits shape: (dim) pub fn simple_sample( logits: &Tensor, @@ -68,7 +79,6 @@ pub fn simple_sample( 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")); @@ -90,7 +100,7 @@ pub fn simple_sample( } if let Some(top_k) = top_k && top_k > 0 - && top_k > logits.dim(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)?; @@ -114,7 +124,7 @@ pub fn simple_sample( .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)?; + mask = mask.slice_scatter(&Tensor::new(&[0u8], logits.device())?, 0, 0)?; } let on_true = Tensor::new(f32::NEG_INFINITY, logits.device())? .to_dtype(logits.dtype())? @@ -122,19 +132,9 @@ pub fn simple_sample( let new_logits = mask.where_cond(&on_true, &sorted_logits)?; logits = logits.scatter(&sorted_indices, &new_logits, 0)?; } + let probs = softmax(&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) + let probs = probs.to_dtype(candle_core::DType::F32)?.to_vec1::()?; + sample_weighted(&probs) } } diff --git a/src/models/glm_asr_nano/model.rs b/src/models/glm_asr_nano/model.rs index ab3f2c4..048baa5 100644 --- a/src/models/glm_asr_nano/model.rs +++ b/src/models/glm_asr_nano/model.rs @@ -6,11 +6,10 @@ use crate::{ models::{ common::{ InferenceModel, - modules::{ - TwoLinearMLP, eager_attention_forward, get_conv1d, get_layer_norm, - }, + modules::{TwoLinearMLP, eager_attention_forward, get_conv1d, get_layer_norm}, }, - glm_asr_nano::config::{GlmAsrAudioConfig, GlmAsrNanoConfig}, llama::LlamaForCausalLM, + 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/gpt2/mod.rs b/src/models/gpt2/mod.rs index 261c5c5..91b8be0 100644 --- a/src/models/gpt2/mod.rs +++ b/src/models/gpt2/mod.rs @@ -71,7 +71,7 @@ impl GPT2Attention { } }; self.kv_cache = Some((key_states.clone(), value_states.clone())); - let scale = 1f64 / f64::sqrt(self.head_dim as f64); + let scale: f64 = 1f64 / f64::sqrt(self.head_dim as f64); let attn_output = eager_attention_forward( &query_states, &key_states, @@ -277,7 +277,8 @@ impl GPT2Model { 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())?; + let (cos, sin) = + rope.forward_repeat_interleave(seqlen_offset, seq_len, inputs_embeds.device())?; (Some(cos), Some(sin)) } else { (None, None) diff --git a/src/models/mod.rs b/src/models/mod.rs index 64d369e..1ffaa08 100644 --- a/src/models/mod.rs +++ b/src/models/mod.rs @@ -12,7 +12,8 @@ pub mod lfm2; pub mod lfm2vl; pub mod mask_gct; pub mod minicpm4; -pub mod moss; +pub mod moss_audio_tokenizer_nano; +pub mod moss_tts_nano; pub mod paddleocr_vl; pub mod qwen2; pub mod qwen2_5vl; diff --git a/src/models/moss_audio_tokenizer_nano/config.rs b/src/models/moss_audio_tokenizer_nano/config.rs new file mode 100644 index 0000000..14faef6 --- /dev/null +++ b/src/models/moss_audio_tokenizer_nano/config.rs @@ -0,0 +1,53 @@ +use serde::Deserialize; + +#[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, +} diff --git a/src/models/moss/audio_tokenizer_nano.rs b/src/models/moss_audio_tokenizer_nano/mod.rs similarity index 91% rename from src/models/moss/audio_tokenizer_nano.rs rename to src/models/moss_audio_tokenizer_nano/mod.rs index 2c814ee..e67ba9f 100644 --- a/src/models/moss/audio_tokenizer_nano.rs +++ b/src/models/moss_audio_tokenizer_nano/mod.rs @@ -1,3 +1,4 @@ +pub mod config; use anyhow::{Result, anyhow}; use candle_core::{D, IndexOp, Tensor}; use candle_nn::{Embedding, LayerNorm, Linear, Module, VarBuilder, embedding, linear_no_bias}; @@ -7,7 +8,7 @@ use crate::{ common::modules::{ TwoLinearMLP, WNConv1d, eager_attention_forward, get_layer_norm, l2_normalize, }, - moss::config::{ + moss_audio_tokenizer_nano::config::{ MossAudioTokenizerConfig, MossAudioTokenizerModuleConfig, MossAudioTokenizerQuantizerKwargs, }, @@ -47,7 +48,8 @@ impl MossAudioTokenizerPatchedPretransform { .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)?; + // let out_lengths = (input_lengths * self.patch_size as f64)?; + let out_lengths = input_lengths.affine(self.patch_size as f64, 0.0)?; Ok((x, out_lengths)) } @@ -398,12 +400,21 @@ impl MossAudioTokenizerLFQ { } Ok((z_q, indices)) } + + pub fn decode_code(&self, codec: &Tensor) -> Result { + let mut z_q = self.codebook.forward(codec)?.transpose(1, 2)?; + if let Some(out_proj) = &self.out_proj { + z_q = out_proj.forward(&z_q)?; + } + Ok(z_q) + } } pub struct MossAudioTokenizerResidualLFQ { input_proj: Option, output_proj: Option, quantizers: Vec, + rvq_dim: usize, } impl MossAudioTokenizerResidualLFQ { @@ -454,6 +465,7 @@ impl MossAudioTokenizerResidualLFQ { input_proj, output_proj, quantizers, + rvq_dim: config.rvq_dim, }) } @@ -483,6 +495,23 @@ impl MossAudioTokenizerResidualLFQ { let all_indices = Tensor::stack(&all_indices, 0)?; Ok(all_indices) } + + pub fn decode_codes(&self, codes: &Tensor) -> Result { + let (_, bs, t) = codes.dims3()?; + let mut emb = Tensor::zeros( + (bs, self.rvq_dim, t), + candle_core::DType::F32, + codes.device(), + )?; + for (i, quantizer) in self.quantizers.iter().enumerate() { + let code_i = quantizer.decode_code(&codes.i(i)?)?; + emb = emb.add(&code_i)?; + } + if let Some(output_proj) = &self.output_proj { + emb = output_proj.forward(&emb)?; + } + Ok(emb) + } } pub struct MossAudioTokenizer { @@ -538,7 +567,7 @@ impl MossAudioTokenizer { if cfg.module_type == "PatchedPretransform" && let Some(patch_size) = cfg.patch_size { - let layer = MossAudioTokenizerPatchedPretransform::new(patch_size, true); + let layer = MossAudioTokenizerPatchedPretransform::new(patch_size, false); decoder.push(MossAudioTokenizerModule::PatchedPretransform(layer)); } else if cfg.module_type == "Transformer" { let context_duration = cfg @@ -617,7 +646,7 @@ impl MossAudioTokenizer { } pub fn encode_one(&self, wav: &Tensor) -> Result { - // (channel, audio_len) -> (bs=1, channel, audio_len) + // in: (channel, audio_len) let (c, len) = wav.dims2()?; if c != self.number_channels { return Err(anyhow!( @@ -632,7 +661,7 @@ impl MossAudioTokenizer { Ok(audio_vec[0].clone()) } - pub fn encode_list(&self, wavs: &Vec) -> Result> { + pub fn encode_list(&self, wavs: &[Tensor]) -> Result> { if wavs.is_empty() { return Err(anyhow!( "MossAudioTokenizer encode_list need wavs len > 0, but the wavs is empty" @@ -664,6 +693,27 @@ impl MossAudioTokenizer { 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)?) + self.batch_encode(&input_values, &length_tensor) + } + + pub fn decode_audio_token_ids_to_waveform(&self, audio_token_ids: &Tensor) -> Result { + let decode_codes = audio_token_ids.t()?.unsqueeze(1)?; //(len, nq) -> (nq, len) -> (nq, 1, len) + let len = decode_codes.dim(2)?; + let mut audio = self.quantizer.decode_codes(&decode_codes)?; + let mut audio_length = Tensor::new(&[len as f32], audio.device())?; + for decoder_module in self.decoder.iter() { + (audio, audio_length) = decoder_module.forward(&audio, &audio_length)?; + } + if self.number_channels == 1 || !self.enable_channel_interleave { + Ok(audio) + } else { + let bs = audio.dim(0)?; + Ok(audio + .squeeze(1)? + .contiguous()? + .reshape((bs, (), self.number_channels))? + .transpose(1, 2)? + .contiguous()?) + } } } diff --git a/src/models/moss/config.rs b/src/models/moss_tts_nano/config.rs similarity index 62% rename from src/models/moss/config.rs rename to src/models/moss_tts_nano/config.rs index 197e163..c994999 100644 --- a/src/models/moss/config.rs +++ b/src/models/moss_tts_nano/config.rs @@ -2,59 +2,6 @@ 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, @@ -67,7 +14,7 @@ pub struct MossTTSConfig { 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>, @@ -85,75 +32,71 @@ pub struct MossTTSConfig { 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_tts_nano/generate.rs similarity index 81% rename from src/models/moss/generate.rs rename to src/models/moss_tts_nano/generate.rs index 034774c..050c46f 100644 --- a/src/models/moss/generate.rs +++ b/src/models/moss_tts_nano/generate.rs @@ -1,11 +1,13 @@ use std::collections::HashMap; use crate::{ - models::moss::{ - audio_tokenizer_nano::MossAudioTokenizer, - config::{MossAudioTokenizerConfig, MossTTSConfig}, - processor::MossTTSProcessor, - tts_nano::{MossTTSMode, MossTTSModel}, + models::{ + moss_audio_tokenizer_nano::{MossAudioTokenizer, config::MossAudioTokenizerConfig}, + moss_tts_nano::{ + config::MossTTSConfig, + model::{MossTTSMode, MossTTSModel}, + processor::MossTTSProcessor, + }, }, utils::{find_type_files, get_device, get_dtype}, }; @@ -33,7 +35,7 @@ impl MossTTSGenerate { 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 audio_dtype = get_dtype(dtype, &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)?; @@ -50,12 +52,10 @@ impl MossTTSGenerate { )?; 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); } } @@ -78,18 +78,21 @@ impl MossTTSGenerate { prompt_text: Option<&str>, mode: Option, ) -> Result<()> { - let (mut input_ids, mask) = self.processor.build_inference_input_ids( + let mode = self.processor.resolved_mode( + mode, + prompt_text.is_some(), + prompt_audio_path.is_some(), + )?; + let input_ids = self.processor.build_inference_input_ids( text, prompt_audio_path, prompt_text, - mode, + mode.clone(), &self.audio_tokenizer, &self.text_tokenizer, &self.device, )?; - let _ = self.model.generate(&input_ids, Some(&mask))?; - // println!("input_ids: {}", input_ids); - // println!("mask: {}", mask); + self.model.generate(&input_ids, &self.audio_tokenizer)?; Ok(()) } } diff --git a/src/models/moss/mod.rs b/src/models/moss_tts_nano/mod.rs similarity index 52% rename from src/models/moss/mod.rs rename to src/models/moss_tts_nano/mod.rs index bfe491d..8b1baf7 100644 --- a/src/models/moss/mod.rs +++ b/src/models/moss_tts_nano/mod.rs @@ -1,5 +1,4 @@ -pub mod audio_tokenizer_nano; pub mod config; pub mod generate; +pub mod model; pub mod processor; -pub mod tts_nano; diff --git a/src/models/moss/tts_nano.rs b/src/models/moss_tts_nano/model.rs similarity index 83% rename from src/models/moss/tts_nano.rs rename to src/models/moss_tts_nano/model.rs index 348696b..c66282d 100644 --- a/src/models/moss/tts_nano.rs +++ b/src/models/moss_tts_nano/model.rs @@ -1,9 +1,16 @@ -use crate::models::{common::sample::simple_sample, gpt2::GPT2Model, moss::config::MossTTSConfig}; +use crate::{ + models::{ + common::sample::simple_sample, gpt2::GPT2Model, + moss_audio_tokenizer_nano::MossAudioTokenizer, moss_tts_nano::config::MossTTSConfig, + }, + utils::audio_utils::save_wav, +}; use anyhow::{Result, anyhow}; use candle_core::{D, IndexOp, Tensor}; use candle_nn::{Embedding, Linear, Module, VarBuilder, embedding, linear_no_bias}; +// use candle_transformers::generation::LogitsProcessor; -#[derive(PartialEq, Debug)] +#[derive(PartialEq, Clone, Debug)] pub enum MossTTSMode { Continuation, VoiceClone, @@ -24,6 +31,7 @@ pub struct MossTTSModel { audio_top_k: usize, audio_top_p: f32, audio_repetition_penalty: f32, + // audio_processor: LogitsProcessor, } impl MossTTSModel { @@ -92,6 +100,7 @@ impl MossTTSModel { audio_top_k: 25, audio_top_p: 0.95, audio_repetition_penalty: 1.2, + // audio_processor, }) } @@ -144,15 +153,9 @@ impl MossTTSModel { .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 { + let logits = Tensor::new(&[slot_token_id_logit, end_token_id_logit], logits.device())?; + let token = simple_sample(&logits, true, None, None, None, None, 1.0)?; + if token == 0 { Ok(self.audio_assistant_slot_token_id) } else { Ok(self.audio_end_token_id) @@ -169,31 +172,28 @@ impl MossTTSModel { 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; + pub fn generate( + &mut self, + input_ids: &Tensor, + audio_tokenizer: &MossAudioTokenizer, + ) -> Result<()> { + let sample_len = 100; 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); + for _ in 0..sample_len { 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(); @@ -216,14 +216,11 @@ impl MossTTSModel { .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]) + 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 = self.audio_processor.sample(&channel_logits)?; let channel_token = simple_sample( &channel_logits, true, @@ -232,25 +229,48 @@ impl MossTTSModel { 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( + 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}"); + let waveform = audio_tokenizer + .decode_audio_token_ids_to_waveform(&audio_token_ids)? + .squeeze(0)?; + save_wav( + &waveform, + "./demo.wav", + 2, + audio_tokenizer.sampling_rate as u32, + )?; + Ok(()) + } + + pub fn decode( + &self, + prompt_audio_code: Option<&Tensor>, + audio_tokenizer: &MossAudioTokenizer, + ) -> Result<()> { + if let Some(audio) = prompt_audio_code { + let waveform = audio_tokenizer + .decode_audio_token_ids_to_waveform(audio)? + .squeeze(0)?; + save_wav( + &waveform, + "./demo.wav", + 2, + audio_tokenizer.sampling_rate as u32, + )?; + } Ok(()) } } diff --git a/src/models/moss/processor.rs b/src/models/moss_tts_nano/processor.rs similarity index 91% rename from src/models/moss/processor.rs rename to src/models/moss_tts_nano/processor.rs index 2d87471..84d3a10 100644 --- a/src/models/moss/processor.rs +++ b/src/models/moss_tts_nano/processor.rs @@ -1,6 +1,7 @@ use crate::{ - models::moss::{ - audio_tokenizer_nano::MossAudioTokenizer, config::MossTTSConfig, tts_nano::MossTTSMode, + models::{ + moss_audio_tokenizer_nano::MossAudioTokenizer, + moss_tts_nano::{config::MossTTSConfig, model::MossTTSMode}, }, tokenizer::sentencepiece_encode_vec, utils::{audio_utils::load_audio_with_resample, prepare_tts_text}, @@ -70,7 +71,7 @@ impl MossTTSProcessor { }) } - fn resolved_mode( + pub fn resolved_mode( &self, mode: Option, has_prompt_text: bool, @@ -99,12 +100,11 @@ impl MossTTSProcessor { text: &str, prompt_audio_path: Option<&str>, prompt_text: Option<&str>, - mode: Option, + mode: MossTTSMode, 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())?; + ) -> Result { let audio_code = if let Some(audio_path) = prompt_audio_path { let audio = load_audio_with_resample( audio_path, @@ -142,7 +142,7 @@ impl MossTTSProcessor { 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, + prompt_audio_codes, self.audio_user_slot_token_id, device, )?; @@ -155,9 +155,7 @@ impl MossTTSProcessor { 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)) + Ok(input_ids) } else { let text = if let Some(prompt_text) = prompt_text { prompt_text + text @@ -176,20 +174,17 @@ impl MossTTSProcessor { 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, + 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)) + Ok(input_ids) } } - fn build_audio_prefix_rows( prompt_audio_codes: &Tensor, slot_token_id: u32, @@ -202,7 +197,7 @@ impl MossTTSProcessor { } fn build_text_raw( - token_ids: &Vec, + token_ids: &[u32], audio_pad_token_id: u32, n_vq: usize, device: &Device, diff --git a/src/models/qwen3_asr/generate.rs b/src/models/qwen3_asr/generate.rs index 65bdaab..e6a5744 100644 --- a/src/models/qwen3_asr/generate.rs +++ b/src/models/qwen3_asr/generate.rs @@ -4,7 +4,8 @@ use crate::{ models::common::{ MultiModalData, generate::{GenerationContext, generate_generic_text}, - modules::{AsrResult, VadFrameResult}, sample::get_logit_processor, + 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/utils/audio_utils.rs b/src/utils/audio_utils.rs index c1d952e..3d688e9 100644 --- a/src/utils/audio_utils.rs +++ b/src/utils/audio_utils.rs @@ -647,7 +647,8 @@ pub fn load_audio_with_resample( 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<()> { +/// audio: (1, len) +pub fn save_wav_mono(audio: &Tensor, save_path: &str, sample_rate: u32) -> Result<()> { let spec = hound::WavSpec { channels: 1, sample_rate, @@ -669,6 +670,49 @@ pub fn save_wav(audio: &Tensor, save_path: &str, sample_rate: u32) -> Result<()> Ok(()) } +/// audio: (channel, len) +pub fn save_wav(audio: &Tensor, save_path: &str, channels: usize, sample_rate: u32) -> Result<()> { + if audio.rank() != 2 { + return Err(anyhow!( + "Audio tensor must be 2D (channels, len), got shape {:?}", + audio.dims() + )); + } + + let (num_channels, len) = audio.dims2()?; + if num_channels != channels { + return Err(anyhow!( + "Channel mismatch: Tensor has {} channels, but spec requires {}", + num_channels, + channels + )); + } + let spec = hound::WavSpec { + channels: channels as u16, + sample_rate, + bits_per_sample: 16, + sample_format: hound::SampleFormat::Int, + }; + + let max_val = audio.abs()?.max_all()?.to_scalar::()?; + let ratio = if max_val > 1.0 { + 32767.0 / max_val + } else { + 32767.0 + }; + let audio_vec_2d = audio.to_vec2::()?; + let mut writer = hound::WavWriter::create(save_path, spec).unwrap(); + for i in 0..len { + for ch in 0..num_channels { + let sample_i16 = (audio_vec_2d[ch][i] * ratio).round() as i16; + let sample_i16 = sample_i16.clamp(-32768, 32767); + writer.write_sample(sample_i16)?; + } + } + writer.finalize()?; + Ok(()) +} + pub fn get_audio_wav_u8(audio: &Tensor, sample_rate: u32) -> Result> { let spec = hound::WavSpec { channels: 1, diff --git a/src/utils/mod.rs b/src/utils/mod.rs index 6c3a3a4..baa7184 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -727,7 +727,8 @@ pub fn contains_cjk(text: &str) -> bool { 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 + || (0xac00..=0xd7af).contains(&c) + // Hangul Syllables { return true; } @@ -737,10 +738,10 @@ pub fn contains_cjk(text: &str) -> bool { 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.")) + if normalized_text.is_empty() { + return Err(anyhow!("Text cannot be empty.")); } - normalized_text = normalized_text.replace('\n', " ").replace('\r', " "); + normalized_text = normalized_text.replace(['\n', '\r'], " "); while normalized_text.contains(" ") { normalized_text = normalized_text.replace(" ", " "); } @@ -753,22 +754,22 @@ pub fn prepare_tts_text(text: &str) -> Result { return Ok(normalized_text); } - // Non-CJK (English/Western) logic + // 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); - } + if let Some(first_char) = normalized_text.chars().next() + && 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('.'); - } + if let Some(last_char) = normalized_text.chars().last() + && last_char.is_alphanumeric() + { + normalized_text.push('.'); } // Add padding if less than 5 words diff --git a/tests/config_tests.rs b/tests/config_tests.rs index 5e7b202..2dc9adf 100644 --- a/tests/config_tests.rs +++ b/tests/config_tests.rs @@ -1,5 +1,5 @@ use aha::models::{ - 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 + deepseek_ocr::config::DeepseekOCRConfig, hunyuan_ocr::config::HunYuanVLConfig, lfm2::config::Lfm2Config, lfm2vl::config::{Lfm2ProcessorConfig, Lfm2VLConfig}, minicpm4::config::MiniCPM4Config, moss_audio_tokenizer_nano::config::MossAudioTokenizerConfig, moss_tts_nano::config::MossTTSConfig, paddleocr_vl::config::PaddleOCRVLConfig, qwen2_5vl::config::Qwen2_5VLConfig, qwen3vl::config::Qwen3VLConfig, voxcpm::config::VoxCPMConfig }; use anyhow::Result; @@ -128,4 +128,4 @@ fn moss_tts_config() -> Result<()> { 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 index 9964783..db630b3 100644 --- a/tests/test_moss_tts.rs +++ b/tests/test_moss_tts.rs @@ -1,4 +1,4 @@ -use aha::models::moss::generate::MossTTSGenerate; +use aha::models::moss_tts_nano::generate::MossTTSGenerate; use anyhow::Result; #[test] @@ -10,11 +10,13 @@ fn moss_tts() -> Result<()> { 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, + Some("file://./assets/audio/jiangjiang.wav"), + None, + // Some("哈喽大家好,我是蒋蒋"), + // Some(aha::models::moss_tts_nano::model::MossTTSMode::Continuation), + None, )?; Ok(()) } diff --git a/tests/test_voxcpm.rs b/tests/test_voxcpm.rs index dd227dd..570e830 100644 --- a/tests/test_voxcpm.rs +++ b/tests/test_voxcpm.rs @@ -6,7 +6,7 @@ use aha::{ GenerateModel, voxcpm::{generate::VoxCPMGenerate, tokenizer::SingleChineseTokenizer}, }, - utils::audio_utils::{extract_and_save_audio_from_response, save_wav}, + utils::audio_utils::{extract_and_save_audio_from_response, save_wav_mono}, }; use anyhow::{Ok, Result}; @@ -54,7 +54,7 @@ fn voxcpm_use_message_generate() -> Result<()> { } let i_duration = i_start.elapsed(); println!("Time elapsed in generate is: {:?}", i_duration); - // save_wav(&generate, "voxcpm.wav", 16000)?; + // save_wav_mono(&generate, "voxcpm.wav", 16000)?; Ok(()) } @@ -105,7 +105,7 @@ fn voxcpm_generate() -> Result<()> { let i_duration = i_start.elapsed(); println!("Time elapsed in generate is: {:?}", i_duration); - save_wav(&generate, "voxcpm.wav", 16000)?; + save_wav_mono(&generate, "voxcpm.wav", 16000)?; Ok(()) } diff --git a/tests/test_voxcpm1_5.rs b/tests/test_voxcpm1_5.rs index 5b5fd33..f3355ba 100644 --- a/tests/test_voxcpm1_5.rs +++ b/tests/test_voxcpm1_5.rs @@ -7,7 +7,7 @@ use aha::{ GenerateModel, voxcpm::{generate::VoxCPMGenerate, tokenizer::SingleChineseTokenizer}, }, - utils::audio_utils::{extract_and_save_audio_from_response, save_wav}, + utils::audio_utils::{extract_and_save_audio_from_response, save_wav_mono}, }; use anyhow::{Ok, Result}; @@ -106,7 +106,7 @@ fn voxcpm1_5_generate() -> Result<()> { let i_duration = i_start.elapsed(); println!("Time elapsed in generate is: {:?}", i_duration); - save_wav(&generate, "voxcpm1_5.wav", 44100)?; + save_wav_mono(&generate, "voxcpm1_5.wav", 44100)?; Ok(()) } @@ -169,7 +169,7 @@ fn voxcpm_refact_generate() -> Result<()> { std::thread::sleep(std::time::Duration::from_secs(2)); let i_duration = i_start.elapsed(); println!("Time elapsed in generate is: {:?}", i_duration); - save_wav( + save_wav_mono( &generate, "voxcpm.wav", voxcpm_generate.sample_rate() as u32, diff --git a/tests/weight_test.rs b/tests/weight_test.rs index 9e80a3f..a73a742 100644 --- a/tests/weight_test.rs +++ b/tests/weight_test.rs @@ -444,11 +444,14 @@ 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 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 +}