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]]
<<<<<<< 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"
+1 -1
View File
@@ -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);
+5 -2
View File
@@ -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::{
+13
View File
@@ -1393,3 +1393,16 @@ pub fn quick_gelu(xs: &Tensor) -> Result<Tensor> {
let x = sigmoid(&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_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<f32>,
@@ -43,7 +43,7 @@ pub fn use_repeat_penalty(
logits: &Tensor,
context: &[u32],
) -> 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())
} 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<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)
pub fn simple_sample(
logits: &Tensor,
@@ -68,7 +79,6 @@ pub fn simple_sample(
top_p: Option<f32>,
previous_token_ids: Option<&[u32]>,
repeat_penalty: f32,
seed: Option<u64>,
) -> Result<u32> {
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::<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())?
.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::<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)
let probs = probs.to_dtype(candle_core::DType::F32)?.to_vec1::<f32>()?;
sample_weighted(&probs)
}
}
+3 -4
View File
@@ -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},
+3 -2
View File
@@ -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<Tensor> {
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)
+2 -1
View File
@@ -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;
@@ -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 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<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 {
input_proj: Option<WNConv1d>,
output_proj: Option<WNConv1d>,
quantizers: Vec<MossAudioTokenizerLFQ>,
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<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 {
@@ -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<Tensor> {
// (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<Tensor>) -> Result<Vec<Tensor>> {
pub fn encode_list(&self, wavs: &[Tensor]) -> Result<Vec<Tensor>> {
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<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;
#[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)]
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<Vec<u32>>,
pub begin_suppress_tokens: Option<Vec<u32>>,
@@ -85,75 +32,71 @@ pub struct MossTTSConfig {
pub finetuning_task: Option<String>,
pub forced_bos_token_id: Option<u32>,
pub forced_eos_token_id: Option<u32>,
// GPT2 Backbone Config
pub gpt2_config: GPT2Config,
pub hidden_size: usize,
pub id2label: std::collections::HashMap<usize, String>,
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<String, usize>,
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<String>,
pub problem_type: Option<String>,
// pub pruned_heads: std::collections::HashMap<String, Vec<usize>>,
pub remove_invalid_values: bool,
pub repetition_penalty: f64,
pub return_dict: bool,
pub return_dict_in_generate: bool,
pub sep_token_id: Option<u32>,
pub suppress_tokens: Option<Vec<u32>>,
pub task_specific_params: Option<serde_json::Value>,
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,
}
@@ -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<MossTTSMode>,
) -> 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(())
}
}
@@ -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;
@@ -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::<f32>()?;
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(&current_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(&current_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(())
}
}
@@ -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<MossTTSMode>,
has_prompt_text: bool,
@@ -99,12 +100,11 @@ impl MossTTSProcessor {
text: &str,
prompt_audio_path: Option<&str>,
prompt_text: Option<&str>,
mode: Option<MossTTSMode>,
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<Tensor> {
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<u32>,
token_ids: &[u32],
audio_pad_token_id: u32,
n_vq: usize,
device: &Device,
+2 -1
View File
@@ -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},
+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)
}
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::<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>> {
let spec = hound::WavSpec {
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
|| (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<String> {
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<String> {
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
+2 -2
View File
@@ -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(())
}
}
+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;
#[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(())
}
+3 -3
View File
@@ -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(())
}
+3 -3
View File
@@ -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,
+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
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(())
}
}