add moss-tts-nano

This commit is contained in:
jhqxxx
2026-05-11 16:15:09 +08:00
parent d38ed80ab4
commit a14b9e1e90
23 changed files with 338 additions and 225 deletions
Generated
-18
View File
@@ -885,7 +885,6 @@ dependencies = [
] ]
[[package]] [[package]]
<<<<<<< HEAD
name = "cpufeatures" name = "cpufeatures"
version = "0.3.0" version = "0.3.0"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
@@ -895,23 +894,6 @@ dependencies = [
] ]
[[package]] [[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" name = "crc32fast"
version = "1.5.0" version = "1.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
+1 -1
View File
@@ -51,7 +51,7 @@ impl ExecModel for VoxCPMExec {
}; };
let sample_rate = voxcpm_generate.sample_rate(); 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); println!("Output saved to: {}", output_path);
+5 -2
View File
@@ -1,12 +1,15 @@
use anyhow::Result; use anyhow::Result;
use candle_core::{DType, Device, Tensor}; use candle_core::{DType, Device, Tensor};
use candle_transformers::generation::{LogitsProcessor}; use candle_transformers::generation::LogitsProcessor;
use rocket::async_stream::stream; use rocket::async_stream::stream;
use rocket::futures::Stream; use rocket::futures::Stream;
use std::time::Instant; use std::time::Instant;
use crate::{ 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}, params::chat::{ChatCompletionChunkResponse, ChatCompletionResponse},
tokenizer::TokenizerModel, tokenizer::TokenizerModel,
utils::response_utils::{ utils::response_utils::{
+13
View File
@@ -1393,3 +1393,16 @@ pub fn quick_gelu(xs: &Tensor) -> Result<Tensor> {
let x = sigmoid(&x)?; let x = sigmoid(&x)?;
Ok(xs.mul(&x)?) Ok(xs.mul(&x)?)
} }
pub fn new_gelu(xs: &Tensor) -> Result<Tensor> {
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)?)
}
+19 -19
View File
@@ -2,7 +2,7 @@ use anyhow::{Result, anyhow};
use candle_core::{IndexOp, Tensor}; use candle_core::{IndexOp, Tensor};
use candle_nn::ops::softmax; use candle_nn::ops::softmax;
use candle_transformers::generation::{LogitsProcessor, Sampling}; use candle_transformers::generation::{LogitsProcessor, Sampling};
use rand::{SeedableRng, distr::Distribution}; use rand::distr::Distribution;
pub fn get_logit_processor( pub fn get_logit_processor(
temperature: Option<f32>, temperature: Option<f32>,
@@ -43,7 +43,7 @@ pub fn use_repeat_penalty(
logits: &Tensor, logits: &Tensor,
context: &[u32], context: &[u32],
) -> Result<Tensor> { ) -> Result<Tensor> {
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()) Ok(logits.clone())
} else { } else {
let start_at = if let Some(last_n) = repeat_last_n { let start_at = if let Some(last_n) = repeat_last_n {
@@ -52,13 +52,24 @@ pub fn use_repeat_penalty(
0 0
}; };
Ok(candle_transformers::utils::apply_repeat_penalty( Ok(candle_transformers::utils::apply_repeat_penalty(
&logits, logits,
repeat_penalty, repeat_penalty,
&context[start_at..], &context[start_at..],
)?) )?)
} }
} }
pub fn sample_weighted(prs: &[f32]) -> Result<u32> {
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) /// logits shape: (dim)
pub fn simple_sample( pub fn simple_sample(
logits: &Tensor, logits: &Tensor,
@@ -68,7 +79,6 @@ pub fn simple_sample(
top_p: Option<f32>, top_p: Option<f32>,
previous_token_ids: Option<&[u32]>, previous_token_ids: Option<&[u32]>,
repeat_penalty: f32, repeat_penalty: f32,
seed: Option<u64>,
) -> Result<u32> { ) -> Result<u32> {
if logits.rank() != 1 { if logits.rank() != 1 {
return Err(anyhow!("simple_sample logits need 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 if let Some(top_k) = top_k
&& top_k > 0 && top_k > 0
&& top_k > logits.dim(0)? && top_k < logits.dim(0)?
{ {
let sorted_indices = logits.arg_sort_last_dim(false)?; let sorted_indices = logits.arg_sort_last_dim(false)?;
let top_k_indices = sorted_indices.narrow(0, 0, top_k)?; 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())?)?; .broadcast_gt(&Tensor::new(top_p, logits.device())?.to_dtype(logits.dtype())?)?;
// 保证数据不会被全部置为-inf // 保证数据不会被全部置为-inf
if mask.i(0)?.to_scalar::<u8>()? == 1 { if mask.i(0)?.to_scalar::<u8>()? == 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())? let on_true = Tensor::new(f32::NEG_INFINITY, logits.device())?
.to_dtype(logits.dtype())? .to_dtype(logits.dtype())?
@@ -122,19 +132,9 @@ pub fn simple_sample(
let new_logits = mask.where_cond(&on_true, &sorted_logits)?; let new_logits = mask.where_cond(&on_true, &sorted_logits)?;
logits = logits.scatter(&sorted_indices, &new_logits, 0)?; logits = logits.scatter(&sorted_indices, &new_logits, 0)?;
} }
let probs = softmax(&logits, 0)?;
let probs = softmax(&logits, 0)? let probs = probs.to_dtype(candle_core::DType::F32)?.to_vec1::<f32>()?;
.to_dtype(candle_core::DType::F32)? sample_weighted(&probs)
.to_vec1::<f32>()?;
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)
} }
} }
+3 -4
View File
@@ -6,11 +6,10 @@ use crate::{
models::{ models::{
common::{ common::{
InferenceModel, InferenceModel,
modules::{ modules::{TwoLinearMLP, eager_attention_forward, get_conv1d, get_layer_norm},
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}, position_embed::rope::{RoPE, glm_asr_apply_rotary_pos_emb},
utils::tensor_utils::{get_equal_mask, masked_scatter_dim0}, utils::tensor_utils::{get_equal_mask, masked_scatter_dim0},
+3 -2
View File
@@ -71,7 +71,7 @@ impl GPT2Attention {
} }
}; };
self.kv_cache = Some((key_states.clone(), value_states.clone())); 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( let attn_output = eager_attention_forward(
&query_states, &query_states,
&key_states, &key_states,
@@ -277,7 +277,8 @@ impl GPT2Model {
pub fn forward(&mut self, inputs_embeds: &Tensor, seqlen_offset: usize) -> Result<Tensor> { pub fn forward(&mut self, inputs_embeds: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
let (b_size, seq_len, _) = inputs_embeds.dims3()?; let (b_size, seq_len, _) = inputs_embeds.dims3()?;
let (cos, sin) = if let Some(rope) = &self.rope { 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)) (Some(cos), Some(sin))
} else { } else {
(None, None) (None, None)
+2 -1
View File
@@ -12,7 +12,8 @@ pub mod lfm2;
pub mod lfm2vl; pub mod lfm2vl;
pub mod mask_gct; pub mod mask_gct;
pub mod minicpm4; pub mod minicpm4;
pub mod moss; pub mod moss_audio_tokenizer_nano;
pub mod moss_tts_nano;
pub mod paddleocr_vl; pub mod paddleocr_vl;
pub mod qwen2; pub mod qwen2;
pub mod qwen2_5vl; pub mod qwen2_5vl;
@@ -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<MossAudioTokenizerModuleConfig>,
pub decoder_kwargs: Vec<MossAudioTokenizerModuleConfig>,
pub quantizer_type: String,
pub quantizer_kwargs: MossAudioTokenizerQuantizerKwargs,
pub reversed_decoder_kwargs: Vec<MossAudioTokenizerModuleConfig>,
}
#[derive(Debug, Deserialize)]
pub struct MossAudioTokenizerModuleConfig {
pub module_type: String,
pub patch_size: Option<usize>,
pub causal: Option<bool>,
pub context_duration: Option<f64>,
pub conv_layout: Option<bool>,
pub d_model: Option<usize>,
pub dim_feedforward: Option<usize>,
pub gating: Option<String>,
pub input_dimension: Option<usize>,
pub layer_scale: Option<f64>,
pub max_period: Option<usize>,
pub norm: Option<String>,
pub num_heads: Option<usize>,
pub num_layers: Option<usize>,
pub output_dimension: Option<usize>,
pub positional_embedding: Option<String>,
}
#[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,
}
@@ -1,3 +1,4 @@
pub mod config;
use anyhow::{Result, anyhow}; use anyhow::{Result, anyhow};
use candle_core::{D, IndexOp, Tensor}; use candle_core::{D, IndexOp, Tensor};
use candle_nn::{Embedding, LayerNorm, Linear, Module, VarBuilder, embedding, linear_no_bias}; use candle_nn::{Embedding, LayerNorm, Linear, Module, VarBuilder, embedding, linear_no_bias};
@@ -7,7 +8,7 @@ use crate::{
common::modules::{ common::modules::{
TwoLinearMLP, WNConv1d, eager_attention_forward, get_layer_norm, l2_normalize, TwoLinearMLP, WNConv1d, eager_attention_forward, get_layer_norm, l2_normalize,
}, },
moss::config::{ moss_audio_tokenizer_nano::config::{
MossAudioTokenizerConfig, MossAudioTokenizerModuleConfig, MossAudioTokenizerConfig, MossAudioTokenizerModuleConfig,
MossAudioTokenizerQuantizerKwargs, MossAudioTokenizerQuantizerKwargs,
}, },
@@ -47,7 +48,8 @@ impl MossAudioTokenizerPatchedPretransform {
.reshape((b, d, self.patch_size, l))? .reshape((b, d, self.patch_size, l))?
.permute((0, 1, 3, 2))? .permute((0, 1, 3, 2))?
.reshape((b, d, l * self.patch_size))?; .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)) Ok((x, out_lengths))
} }
@@ -398,12 +400,21 @@ impl MossAudioTokenizerLFQ {
} }
Ok((z_q, indices)) Ok((z_q, indices))
} }
pub fn decode_code(&self, codec: &Tensor) -> Result<Tensor> {
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 { pub struct MossAudioTokenizerResidualLFQ {
input_proj: Option<WNConv1d>, input_proj: Option<WNConv1d>,
output_proj: Option<WNConv1d>, output_proj: Option<WNConv1d>,
quantizers: Vec<MossAudioTokenizerLFQ>, quantizers: Vec<MossAudioTokenizerLFQ>,
rvq_dim: usize,
} }
impl MossAudioTokenizerResidualLFQ { impl MossAudioTokenizerResidualLFQ {
@@ -454,6 +465,7 @@ impl MossAudioTokenizerResidualLFQ {
input_proj, input_proj,
output_proj, output_proj,
quantizers, quantizers,
rvq_dim: config.rvq_dim,
}) })
} }
@@ -483,6 +495,23 @@ impl MossAudioTokenizerResidualLFQ {
let all_indices = Tensor::stack(&all_indices, 0)?; let all_indices = Tensor::stack(&all_indices, 0)?;
Ok(all_indices) Ok(all_indices)
} }
pub fn decode_codes(&self, codes: &Tensor) -> Result<Tensor> {
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 { pub struct MossAudioTokenizer {
@@ -538,7 +567,7 @@ impl MossAudioTokenizer {
if cfg.module_type == "PatchedPretransform" if cfg.module_type == "PatchedPretransform"
&& let Some(patch_size) = cfg.patch_size && 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)); decoder.push(MossAudioTokenizerModule::PatchedPretransform(layer));
} else if cfg.module_type == "Transformer" { } else if cfg.module_type == "Transformer" {
let context_duration = cfg let context_duration = cfg
@@ -617,7 +646,7 @@ impl MossAudioTokenizer {
} }
pub fn encode_one(&self, wav: &Tensor) -> Result<Tensor> { pub fn encode_one(&self, wav: &Tensor) -> Result<Tensor> {
// (channel, audio_len) -> (bs=1, channel, audio_len) // in: (channel, audio_len)
let (c, len) = wav.dims2()?; let (c, len) = wav.dims2()?;
if c != self.number_channels { if c != self.number_channels {
return Err(anyhow!( return Err(anyhow!(
@@ -632,7 +661,7 @@ impl MossAudioTokenizer {
Ok(audio_vec[0].clone()) Ok(audio_vec[0].clone())
} }
pub fn encode_list(&self, wavs: &Vec<Tensor>) -> Result<Vec<Tensor>> { pub fn encode_list(&self, wavs: &[Tensor]) -> Result<Vec<Tensor>> {
if wavs.is_empty() { if wavs.is_empty() {
return Err(anyhow!( return Err(anyhow!(
"MossAudioTokenizer encode_list need wavs len > 0, but the wavs is empty" "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 input_values = Tensor::stack(&input_values, 0)?;
let length_tensor = Tensor::new(length.clone(), input_values.device())? let length_tensor = Tensor::new(length.clone(), input_values.device())?
.to_dtype(candle_core::DType::F32)?; .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<Tensor> {
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()?)
}
} }
} }
@@ -2,59 +2,6 @@ use serde::Deserialize;
use crate::models::gpt2::config::GPT2Config; 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<MossAudioTokenizerModuleConfig>,
pub decoder_kwargs: Vec<MossAudioTokenizerModuleConfig>,
pub quantizer_type: String,
pub quantizer_kwargs: MossAudioTokenizerQuantizerKwargs,
pub reversed_decoder_kwargs: Vec<MossAudioTokenizerModuleConfig>,
}
#[derive(Debug, Deserialize)]
pub struct MossAudioTokenizerModuleConfig {
pub module_type: String,
pub patch_size: Option<usize>,
pub causal: Option<bool>,
pub context_duration: Option<f64>,
pub conv_layout: Option<bool>,
pub d_model: Option<usize>,
pub dim_feedforward: Option<usize>,
pub gating: Option<String>,
pub input_dimension: Option<usize>,
pub layer_scale: Option<f64>,
pub max_period: Option<usize>,
pub norm: Option<String>,
pub num_heads: Option<usize>,
pub num_layers: Option<usize>,
pub output_dimension: Option<usize>,
pub positional_embedding: Option<String>,
}
#[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)] #[derive(Debug, Deserialize)]
pub struct MossTTSConfig { pub struct MossTTSConfig {
pub add_cross_attention: bool, pub add_cross_attention: bool,
@@ -67,7 +14,7 @@ pub struct MossTTSConfig {
pub audio_tokenizer_sample_rate: usize, pub audio_tokenizer_sample_rate: usize,
pub audio_user_slot_token_id: u32, pub audio_user_slot_token_id: u32,
pub audio_vocab_size: usize, pub audio_vocab_size: usize,
// Generation/Model Params (Simplified nullables to Options or defaults if not critical) // Generation/Model Params (Simplified nullables to Options or defaults if not critical)
pub bad_words_ids: Option<Vec<u32>>, pub bad_words_ids: Option<Vec<u32>>,
pub begin_suppress_tokens: Option<Vec<u32>>, pub begin_suppress_tokens: Option<Vec<u32>>,
@@ -85,75 +32,71 @@ pub struct MossTTSConfig {
pub finetuning_task: Option<String>, pub finetuning_task: Option<String>,
pub forced_bos_token_id: Option<u32>, pub forced_bos_token_id: Option<u32>,
pub forced_eos_token_id: Option<u32>, pub forced_eos_token_id: Option<u32>,
// GPT2 Backbone Config // GPT2 Backbone Config
pub gpt2_config: GPT2Config, pub gpt2_config: GPT2Config,
pub hidden_size: usize, pub hidden_size: usize,
pub id2label: std::collections::HashMap<usize, String>, pub id2label: std::collections::HashMap<usize, String>,
pub im_end_token_id: u32, pub im_end_token_id: u32,
pub im_start_token_id: u32, pub im_start_token_id: u32,
pub initializer_range: f64, pub initializer_range: f64,
pub is_decoder: bool, pub is_decoder: bool,
pub is_encoder_decoder: bool, pub is_encoder_decoder: bool,
pub label2id: std::collections::HashMap<String, usize>, pub label2id: std::collections::HashMap<String, usize>,
pub length_penalty: f64, pub length_penalty: f64,
pub local_transformer_attn_implementation: String, pub local_transformer_attn_implementation: String,
pub local_transformer_layers: usize, pub local_transformer_layers: usize,
pub max_length: usize, pub max_length: usize,
pub max_position_embeddings: usize, pub max_position_embeddings: usize,
pub min_length: usize, pub min_length: usize,
pub model_architecture: String, pub model_architecture: String,
pub model_type: String, pub model_type: String,
pub n_vq: usize, pub n_vq: usize,
pub no_repeat_ngram_size: usize, pub no_repeat_ngram_size: usize,
pub num_beam_groups: usize, pub num_beam_groups: usize,
pub num_beams: usize, pub num_beams: usize,
pub num_return_sequences: usize, pub num_return_sequences: usize,
pub output_attentions: bool, pub output_attentions: bool,
pub output_hidden_states: bool, pub output_hidden_states: bool,
pub output_scores: bool, pub output_scores: bool,
pub pad_token_id: u32, pub pad_token_id: u32,
pub prefix: Option<String>, pub prefix: Option<String>,
pub problem_type: Option<String>, pub problem_type: Option<String>,
// pub pruned_heads: std::collections::HashMap<String, Vec<usize>>, // pub pruned_heads: std::collections::HashMap<String, Vec<usize>>,
pub remove_invalid_values: bool, pub remove_invalid_values: bool,
pub repetition_penalty: f64, pub repetition_penalty: f64,
pub return_dict: bool, pub return_dict: bool,
pub return_dict_in_generate: bool, pub return_dict_in_generate: bool,
pub sep_token_id: Option<u32>, pub sep_token_id: Option<u32>,
pub suppress_tokens: Option<Vec<u32>>, pub suppress_tokens: Option<Vec<u32>>,
pub task_specific_params: Option<serde_json::Value>, pub task_specific_params: Option<serde_json::Value>,
pub temperature: f32, pub temperature: f32,
pub tf_legacy_loss: bool, pub tf_legacy_loss: bool,
pub tie_encoder_decoder: bool, pub tie_encoder_decoder: bool,
pub tie_word_embeddings: bool, pub tie_word_embeddings: bool,
pub tokenizer_class: String, pub tokenizer_class: String,
pub tokenizer_use_fast: bool, pub tokenizer_use_fast: bool,
pub top_k: usize, pub top_k: usize,
pub top_p: f32, pub top_p: f32,
pub torchscript: bool, pub torchscript: bool,
pub typical_p: f64, pub typical_p: f64,
pub use_bfloat16: bool, pub use_bfloat16: bool,
pub vocab_size: usize, pub vocab_size: usize,
} }
@@ -1,11 +1,13 @@
use std::collections::HashMap; use std::collections::HashMap;
use crate::{ use crate::{
models::moss::{ models::{
audio_tokenizer_nano::MossAudioTokenizer, moss_audio_tokenizer_nano::{MossAudioTokenizer, config::MossAudioTokenizerConfig},
config::{MossAudioTokenizerConfig, MossTTSConfig}, moss_tts_nano::{
processor::MossTTSProcessor, config::MossTTSConfig,
tts_nano::{MossTTSMode, MossTTSModel}, model::{MossTTSMode, MossTTSModel},
processor::MossTTSProcessor,
},
}, },
utils::{find_type_files, get_device, get_dtype}, utils::{find_type_files, get_device, get_dtype},
}; };
@@ -33,7 +35,7 @@ impl MossTTSGenerate {
let audio_tokenizer_cfg: MossAudioTokenizerConfig = let audio_tokenizer_cfg: MossAudioTokenizerConfig =
serde_json::from_slice(&std::fs::read(audio_tokenizer_config_path)?)?; serde_json::from_slice(&std::fs::read(audio_tokenizer_config_path)?)?;
let model_list = find_type_files(audio_tokenizer_path, "safetensors")?; 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 device = get_device(device);
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, audio_dtype, &device)? }; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, audio_dtype, &device)? };
let audio_tokenizer = MossAudioTokenizer::new(vb, &audio_tokenizer_cfg)?; 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 model_list = find_type_files(tts_path, "bin")?;
let mut dict_to_hashmap = HashMap::new(); let mut dict_to_hashmap = HashMap::new();
// let cfg_dtype = tts_cfg.dtype.as_str();
let m_dtype = get_dtype(dtype, "bfloat16"); let m_dtype = get_dtype(dtype, "bfloat16");
for m in model_list { for m in model_list {
let dict = read_all_with_key(m, None)?; let dict = read_all_with_key(m, None)?;
for (k, v) in dict { for (k, v) in dict {
// println!("key: {}, tensor shape: {:?}", k, v);
dict_to_hashmap.insert(k, v); dict_to_hashmap.insert(k, v);
} }
} }
@@ -78,18 +78,21 @@ impl MossTTSGenerate {
prompt_text: Option<&str>, prompt_text: Option<&str>,
mode: Option<MossTTSMode>, mode: Option<MossTTSMode>,
) -> Result<()> { ) -> 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, text,
prompt_audio_path, prompt_audio_path,
prompt_text, prompt_text,
mode, mode.clone(),
&self.audio_tokenizer, &self.audio_tokenizer,
&self.text_tokenizer, &self.text_tokenizer,
&self.device, &self.device,
)?; )?;
let _ = self.model.generate(&input_ids, Some(&mask))?; self.model.generate(&input_ids, &self.audio_tokenizer)?;
// println!("input_ids: {}", input_ids);
// println!("mask: {}", mask);
Ok(()) Ok(())
} }
} }
@@ -1,5 +1,4 @@
pub mod audio_tokenizer_nano;
pub mod config; pub mod config;
pub mod generate; pub mod generate;
pub mod model;
pub mod processor; pub mod processor;
pub mod tts_nano;
@@ -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 anyhow::{Result, anyhow};
use candle_core::{D, IndexOp, Tensor}; use candle_core::{D, IndexOp, Tensor};
use candle_nn::{Embedding, Linear, Module, VarBuilder, embedding, linear_no_bias}; 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 { pub enum MossTTSMode {
Continuation, Continuation,
VoiceClone, VoiceClone,
@@ -24,6 +31,7 @@ pub struct MossTTSModel {
audio_top_k: usize, audio_top_k: usize,
audio_top_p: f32, audio_top_p: f32,
audio_repetition_penalty: f32, audio_repetition_penalty: f32,
// audio_processor: LogitsProcessor,
} }
impl MossTTSModel { impl MossTTSModel {
@@ -92,6 +100,7 @@ impl MossTTSModel {
audio_top_k: 25, audio_top_k: 25,
audio_top_p: 0.95, audio_top_p: 0.95,
audio_repetition_penalty: 1.2, audio_repetition_penalty: 1.2,
// audio_processor,
}) })
} }
@@ -144,15 +153,9 @@ impl MossTTSModel {
.i(self.audio_end_token_id)? .i(self.audio_end_token_id)?
.to_dtype(candle_core::DType::F32)? .to_dtype(candle_core::DType::F32)?
.to_scalar::<f32>()?; .to_scalar::<f32>()?;
println!( let logits = Tensor::new(&[slot_token_id_logit, end_token_id_logit], logits.device())?;
"slot_token_id: {} logit: {slot_token_id_logit}", let token = simple_sample(&logits, true, None, None, None, None, 1.0)?;
self.audio_assistant_slot_token_id if token == 0 {
);
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) Ok(self.audio_assistant_slot_token_id)
} else { } else {
Ok(self.audio_end_token_id) Ok(self.audio_end_token_id)
@@ -169,31 +172,28 @@ impl MossTTSModel {
Ok(Tensor::cat(&[&slot, &audio_token_ids], D::Minus1)?) Ok(Tensor::cat(&[&slot, &audio_token_ids], D::Minus1)?)
} }
pub fn generate(&mut self, input_ids: &Tensor, mask: Option<&Tensor>) -> Result<()> { pub fn generate(
let sample_len = 2; &mut self,
input_ids: &Tensor,
audio_tokenizer: &MossAudioTokenizer,
) -> Result<()> {
let sample_len = 100;
let mut seqlen_offset = 0; let mut seqlen_offset = 0;
let mut seq_len = input_ids.dim(1)?; let mut seq_len = input_ids.dim(1)?;
let mut generated_frames = vec![]; let mut generated_frames = vec![];
let mut current_model_input_ids = input_ids.clone(); let mut current_model_input_ids = input_ids.clone();
for step_index in 0..sample_len { for _ in 0..sample_len {
// println!("current_model_input_ids: {:?}", current_model_input_ids);
let inputs_embeds = self.build_inputs_embeds(&current_model_input_ids)?; let inputs_embeds = self.build_inputs_embeds(&current_model_input_ids)?;
let outputs = self.transformer.forward(&inputs_embeds, seqlen_offset)?; let outputs = self.transformer.forward(&inputs_embeds, seqlen_offset)?;
// println!("transformer-----------------------");
let outputs_len = outputs.dim(1)?; let outputs_len = outputs.dim(1)?;
let global_hidden_state = outputs.narrow(1, outputs_len - 1, 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 mut local_positions = 0usize;
let local_outputs = self let local_outputs = self
.local_transformer .local_transformer
.forward(&global_hidden_state, local_positions)?; .forward(&global_hidden_state, local_positions)?;
// println!("local_outputs-----------------------");
let local_len = local_outputs.dim(1)?; let local_len = local_outputs.dim(1)?;
let local_hidden_states = local_outputs.narrow(1, local_len - 1, 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)?; 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)?; let next_text_token = self.sample_next_assistant_text_token(&text_logits)?;
if next_text_token == self.audio_end_token_id { if next_text_token == self.audio_end_token_id {
self.local_transformer.clear_kv_cache(); self.local_transformer.clear_kv_cache();
@@ -216,14 +216,11 @@ impl MossTTSModel {
.forward(&current_local_input, local_positions)?; .forward(&current_local_input, local_positions)?;
let local_len = local_outputs.dim(1)?; let local_len = local_outputs.dim(1)?;
let local_hidden_states = local_outputs.narrow(1, local_len - 1, 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)? .forward(&local_hidden_states)?
.squeeze(0)? .squeeze(0)?
.squeeze(0)?; .squeeze(0)?;
// println!("channel_logits: {}", channel_logits.i(0..100)?); // let channel_token = self.audio_processor.sample(&channel_logits)?;
let arg_max = channel_logits.argmax(0)?;
println!("arg_max: {}", arg_max);
let channel_token = simple_sample( let channel_token = simple_sample(
&channel_logits, &channel_logits,
true, true,
@@ -232,25 +229,48 @@ impl MossTTSModel {
Some(self.audio_top_p), Some(self.audio_top_p),
Some(&next_frame_tokens), Some(&next_frame_tokens),
self.audio_repetition_penalty, self.audio_repetition_penalty,
None,
)?; )?;
println!("channel_token: {channel_token}");
next_frame_tokens.push(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())?, &Tensor::from_slice(&[channel_token], (1, 1), input_ids.device())?,
)?; )?;
// println!("current_local_input: {current_local_input}");
} }
self.local_transformer.clear_kv_cache(); self.local_transformer.clear_kv_cache();
let next_frame = Tensor::new(next_frame_tokens, input_ids.device())?; 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)?; current_model_input_ids = self.build_generation_row(&next_frame)?;
seqlen_offset += seq_len; seqlen_offset += seq_len;
seq_len = 1; seq_len = 1;
generated_frames.push(next_frame); generated_frames.push(next_frame);
} }
let audio_token_ids = Tensor::stack(&generated_frames, 0)?; 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(()) Ok(())
} }
} }
@@ -1,6 +1,7 @@
use crate::{ use crate::{
models::moss::{ models::{
audio_tokenizer_nano::MossAudioTokenizer, config::MossTTSConfig, tts_nano::MossTTSMode, moss_audio_tokenizer_nano::MossAudioTokenizer,
moss_tts_nano::{config::MossTTSConfig, model::MossTTSMode},
}, },
tokenizer::sentencepiece_encode_vec, tokenizer::sentencepiece_encode_vec,
utils::{audio_utils::load_audio_with_resample, prepare_tts_text}, utils::{audio_utils::load_audio_with_resample, prepare_tts_text},
@@ -70,7 +71,7 @@ impl MossTTSProcessor {
}) })
} }
fn resolved_mode( pub fn resolved_mode(
&self, &self,
mode: Option<MossTTSMode>, mode: Option<MossTTSMode>,
has_prompt_text: bool, has_prompt_text: bool,
@@ -99,12 +100,11 @@ impl MossTTSProcessor {
text: &str, text: &str,
prompt_audio_path: Option<&str>, prompt_audio_path: Option<&str>,
prompt_text: Option<&str>, prompt_text: Option<&str>,
mode: Option<MossTTSMode>, mode: MossTTSMode,
audio_tokenizer: &MossAudioTokenizer, audio_tokenizer: &MossAudioTokenizer,
text_tokenizer: &SentencePieceProcessor, text_tokenizer: &SentencePieceProcessor,
device: &Device, device: &Device,
) -> Result<(Tensor, Tensor)> { ) -> Result<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_code = if let Some(audio_path) = prompt_audio_path {
let audio = load_audio_with_resample( let audio = load_audio_with_resample(
audio_path, audio_path,
@@ -142,7 +142,7 @@ impl MossTTSProcessor {
suffix_token_ids.extend_from_slice(&self.assistant_ids); suffix_token_ids.extend_from_slice(&self.assistant_ids);
suffix_token_ids.push(self.audio_start_token_id); suffix_token_ids.push(self.audio_start_token_id);
let audio_prefix_rows = Self::build_audio_prefix_rows( let audio_prefix_rows = Self::build_audio_prefix_rows(
&prompt_audio_codes, prompt_audio_codes,
self.audio_user_slot_token_id, self.audio_user_slot_token_id,
device, device,
)?; )?;
@@ -155,9 +155,7 @@ impl MossTTSProcessor {
let input_ids = let input_ids =
Tensor::cat(&[&prompt_ids_tensor, &audio_prefix_rows, &suffix_rows], 0)? Tensor::cat(&[&prompt_ids_tensor, &audio_prefix_rows, &suffix_rows], 0)?
.unsqueeze(0)?; .unsqueeze(0)?;
let (bs, len, _) = input_ids.dims3()?; Ok(input_ids)
let mask = Tensor::ones((bs, len), candle_core::DType::U32, device)?;
Ok((input_ids, mask))
} else { } else {
let text = if let Some(prompt_text) = prompt_text { let text = if let Some(prompt_text) = prompt_text {
prompt_text + 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)?; Self::build_text_raw(&prompt_ids, self.audio_pad_token_id, self.n_vq, device)?;
if let Some(prompt_audio_codes) = &audio_code { if let Some(prompt_audio_codes) = &audio_code {
let audio_prefix_rows = Self::build_audio_prefix_rows( let audio_prefix_rows = Self::build_audio_prefix_rows(
&prompt_audio_codes, prompt_audio_codes,
self.audio_assistant_slot_token_id, self.audio_assistant_slot_token_id,
device, device,
)?; )?;
input_ids = Tensor::cat(&[&input_ids, &audio_prefix_rows], 0)?; input_ids = Tensor::cat(&[&input_ids, &audio_prefix_rows], 0)?;
} }
input_ids = input_ids.unsqueeze(0)?; input_ids = input_ids.unsqueeze(0)?;
let (bs, len, _) = input_ids.dims3()?; Ok(input_ids)
let mask = Tensor::ones((bs, len), candle_core::DType::U32, device)?;
Ok((input_ids, mask))
} }
} }
fn build_audio_prefix_rows( fn build_audio_prefix_rows(
prompt_audio_codes: &Tensor, prompt_audio_codes: &Tensor,
slot_token_id: u32, slot_token_id: u32,
@@ -202,7 +197,7 @@ impl MossTTSProcessor {
} }
fn build_text_raw( fn build_text_raw(
token_ids: &Vec<u32>, token_ids: &[u32],
audio_pad_token_id: u32, audio_pad_token_id: u32,
n_vq: usize, n_vq: usize,
device: &Device, device: &Device,
+2 -1
View File
@@ -4,7 +4,8 @@ use crate::{
models::common::{ models::common::{
MultiModalData, MultiModalData,
generate::{GenerationContext, generate_generic_text}, generate::{GenerationContext, generate_generic_text},
modules::{AsrResult, VadFrameResult}, sample::get_logit_processor, modules::{AsrResult, VadFrameResult},
sample::get_logit_processor,
}, },
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}, params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
utils::response_utils::{build_chunk_response_with_usage, build_completion_response_with_time}, utils::response_utils::{build_chunk_response_with_usage, build_completion_response_with_time},
+45 -1
View File
@@ -647,7 +647,8 @@ pub fn load_audio_with_resample(
resample_audio_from_bytes(audio_vec, device, target_sample_rate, target_channels) 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 { let spec = hound::WavSpec {
channels: 1, channels: 1,
sample_rate, sample_rate,
@@ -669,6 +670,49 @@ pub fn save_wav(audio: &Tensor, save_path: &str, sample_rate: u32) -> Result<()>
Ok(()) 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::<f32>()?;
let ratio = if max_val > 1.0 {
32767.0 / max_val
} else {
32767.0
};
let audio_vec_2d = audio.to_vec2::<f32>()?;
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<Vec<u8>> { pub fn get_audio_wav_u8(audio: &Tensor, sample_rate: u32) -> Result<Vec<u8>> {
let spec = hound::WavSpec { let spec = hound::WavSpec {
channels: 1, channels: 1,
+17 -16
View File
@@ -727,7 +727,8 @@ pub fn contains_cjk(text: &str) -> bool {
if (0x4e00..=0x9fff).contains(&c) // CJK Unified Ideographs if (0x4e00..=0x9fff).contains(&c) // CJK Unified Ideographs
|| (0x3400..=0x4dbf).contains(&c) // CJK Unified Ideographs Extension A || (0x3400..=0x4dbf).contains(&c) // CJK Unified Ideographs Extension A
|| (0x3040..=0x30ff).contains(&c) // Hiragana and Katakana || (0x3040..=0x30ff).contains(&c) // Hiragana and Katakana
|| (0xac00..=0xd7af).contains(&c) // Hangul Syllables || (0xac00..=0xd7af).contains(&c)
// Hangul Syllables
{ {
return true; return true;
} }
@@ -737,10 +738,10 @@ pub fn contains_cjk(text: &str) -> bool {
pub fn prepare_tts_text(text: &str) -> Result<String> { pub fn prepare_tts_text(text: &str) -> Result<String> {
let mut normalized_text = text.trim().to_string(); let mut normalized_text = text.trim().to_string();
if normalized_text.eq("") { if normalized_text.is_empty() {
return Err(anyhow!("Text cannot be 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(" ") { while normalized_text.contains(" ") {
normalized_text = normalized_text.replace(" ", " "); normalized_text = normalized_text.replace(" ", " ");
} }
@@ -753,22 +754,22 @@ pub fn prepare_tts_text(text: &str) -> Result<String> {
return Ok(normalized_text); return Ok(normalized_text);
} }
// Non-CJK (English/Western) logic // Non-CJK (English/Western) logic
// Capitalize first letter if it's lowercase alphabetic // Capitalize first letter if it's lowercase alphabetic
if let Some(first_char) = normalized_text.chars().next() { if let Some(first_char) = normalized_text.chars().next()
if first_char.is_ascii_lowercase() { && first_char.is_ascii_lowercase()
let mut chars = normalized_text.chars(); {
chars.next(); // consume first char let mut chars = normalized_text.chars();
let rest: String = chars.collect(); chars.next(); // consume first char
normalized_text = format!("{}{}", first_char.to_ascii_uppercase(), rest); let rest: String = chars.collect();
} normalized_text = format!("{}{}", first_char.to_ascii_uppercase(), rest);
} }
// Add period if ends with alphanumeric // Add period if ends with alphanumeric
if let Some(last_char) = normalized_text.chars().last() { if let Some(last_char) = normalized_text.chars().last()
if last_char.is_alphanumeric() { && last_char.is_alphanumeric()
normalized_text.push('.'); {
} normalized_text.push('.');
} }
// Add padding if less than 5 words // Add padding if less than 5 words
+2 -2
View File
@@ -1,5 +1,5 @@
use aha::models::{ 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; use anyhow::Result;
@@ -128,4 +128,4 @@ fn moss_tts_config() -> Result<()> {
let config: MossTTSConfig = serde_json::from_slice(&std::fs::read(config_path)?)?; let config: MossTTSConfig = serde_json::from_slice(&std::fs::read(config_path)?)?;
println!("{:?}", config); println!("{:?}", config);
Ok(()) Ok(())
} }
+7 -5
View File
@@ -1,4 +1,4 @@
use aha::models::moss::generate::MossTTSGenerate; use aha::models::moss_tts_nano::generate::MossTTSGenerate;
use anyhow::Result; use anyhow::Result;
#[test] #[test]
@@ -10,11 +10,13 @@ fn moss_tts() -> Result<()> {
let audio_tokenizer_path = format!("{}/openmoss/MOSS-Audio-Tokenizer-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 mut model = MossTTSGenerate::init(&tts_path, &audio_tokenizer_path, None, None)?;
let _ = model.generate( let _ = model.generate(
"您好啊,吃饭了吗,吃的啥啊中午", "你在干吗啊",
Some("file://./assets/audio/jiangjiang.wav"),
Some("哈喽大家好,我是蒋蒋"),
Some(aha::models::moss::tts_nano::MossTTSMode::Continuation),
// None, // None,
Some("file://./assets/audio/jiangjiang.wav"),
None,
// Some("哈喽大家好,我是蒋蒋"),
// Some(aha::models::moss_tts_nano::model::MossTTSMode::Continuation),
None,
)?; )?;
Ok(()) Ok(())
} }
+3 -3
View File
@@ -6,7 +6,7 @@ use aha::{
GenerateModel, GenerateModel,
voxcpm::{generate::VoxCPMGenerate, tokenizer::SingleChineseTokenizer}, 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}; use anyhow::{Ok, Result};
@@ -54,7 +54,7 @@ fn voxcpm_use_message_generate() -> Result<()> {
} }
let i_duration = i_start.elapsed(); let i_duration = i_start.elapsed();
println!("Time elapsed in generate is: {:?}", i_duration); println!("Time elapsed in generate is: {:?}", i_duration);
// save_wav(&generate, "voxcpm.wav", 16000)?; // save_wav_mono(&generate, "voxcpm.wav", 16000)?;
Ok(()) Ok(())
} }
@@ -105,7 +105,7 @@ fn voxcpm_generate() -> Result<()> {
let i_duration = i_start.elapsed(); let i_duration = i_start.elapsed();
println!("Time elapsed in generate is: {:?}", i_duration); println!("Time elapsed in generate is: {:?}", i_duration);
save_wav(&generate, "voxcpm.wav", 16000)?; save_wav_mono(&generate, "voxcpm.wav", 16000)?;
Ok(()) Ok(())
} }
+3 -3
View File
@@ -7,7 +7,7 @@ use aha::{
GenerateModel, GenerateModel,
voxcpm::{generate::VoxCPMGenerate, tokenizer::SingleChineseTokenizer}, 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}; use anyhow::{Ok, Result};
@@ -106,7 +106,7 @@ fn voxcpm1_5_generate() -> Result<()> {
let i_duration = i_start.elapsed(); let i_duration = i_start.elapsed();
println!("Time elapsed in generate is: {:?}", i_duration); println!("Time elapsed in generate is: {:?}", i_duration);
save_wav(&generate, "voxcpm1_5.wav", 44100)?; save_wav_mono(&generate, "voxcpm1_5.wav", 44100)?;
Ok(()) Ok(())
} }
@@ -169,7 +169,7 @@ fn voxcpm_refact_generate() -> Result<()> {
std::thread::sleep(std::time::Duration::from_secs(2)); std::thread::sleep(std::time::Duration::from_secs(2));
let i_duration = i_start.elapsed(); let i_duration = i_start.elapsed();
println!("Time elapsed in generate is: {:?}", i_duration); println!("Time elapsed in generate is: {:?}", i_duration);
save_wav( save_wav_mono(
&generate, &generate,
"voxcpm.wav", "voxcpm.wav",
voxcpm_generate.sample_rate() as u32, voxcpm_generate.sample_rate() as u32,
+5 -2
View File
@@ -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 // cargo test -F cuda --test weight_test moss_audio_tokenizer_nano_weight -r -- --nocapture
let save_dir = let save_dir =
aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get 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 device = get_device(None);
let weights = safetensors::load(model_path, &device)?; let weights = safetensors::load(model_path, &device)?;
for (key, tensor) in weights.iter() { for (key, tensor) in weights.iter() {
println!("=== {} === {:?}", key, tensor); println!("=== {} === {:?}", key, tensor);
} }
Ok(()) Ok(())
} }