add moss-tts-nano
This commit is contained in:
+1
-1
@@ -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);
|
||||
|
||||
|
||||
@@ -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::{
|
||||
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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},
|
||||
|
||||
@@ -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
@@ -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,
|
||||
}
|
||||
+56
-6
@@ -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(¤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(())
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
@@ -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},
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user