add moss-tts-nano
This commit is contained in:
Generated
-18
@@ -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
@@ -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);
|
||||||
|
|
||||||
|
|||||||
@@ -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::{
|
||||||
|
|||||||
@@ -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
@@ -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)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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},
|
||||||
|
|||||||
@@ -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
@@ -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,
|
||||||
|
}
|
||||||
+56
-6
@@ -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(¤t_model_input_ids)?;
|
let inputs_embeds = self.build_inputs_embeds(¤t_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(¤t_local_input, local_positions)?;
|
.forward(¤t_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,
|
||||||
@@ -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},
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -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(())
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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(())
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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(())
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user