update voxcpm
This commit is contained in:
@@ -179,7 +179,7 @@ impl AttentionNobias {
|
|||||||
Ok(attn_output)
|
Ok(attn_output)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn forward_step(
|
pub fn forward_with_cache(
|
||||||
&mut self,
|
&mut self,
|
||||||
xs: &Tensor,
|
xs: &Tensor,
|
||||||
cos: &Tensor,
|
cos: &Tensor,
|
||||||
|
|||||||
@@ -66,7 +66,7 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
|
|||||||
None => 2048,
|
None => 2048,
|
||||||
};
|
};
|
||||||
for _ in 0..sample_len {
|
for _ in 0..sample_len {
|
||||||
let logits = self.minicpm.forward_step(&input_ids, seqlen_offset)?;
|
let logits = self.minicpm.forward_with_cache(&input_ids, seqlen_offset)?;
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
let next_token = logit_processor.sample(&logits)?;
|
||||||
generate.push(next_token);
|
generate.push(next_token);
|
||||||
@@ -98,7 +98,7 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
|
|||||||
let stream = stream! {
|
let stream = stream! {
|
||||||
let mut error_tokens = Vec::new();
|
let mut error_tokens = Vec::new();
|
||||||
for _ in 0..sample_len {
|
for _ in 0..sample_len {
|
||||||
let logits = self.minicpm.forward_step(
|
let logits = self.minicpm.forward_with_cache(
|
||||||
&input_ids,
|
&input_ids,
|
||||||
seqlen_offset,
|
seqlen_offset,
|
||||||
)?;
|
)?;
|
||||||
|
|||||||
@@ -159,7 +159,7 @@ impl MiniCPMDecoderLayer {
|
|||||||
Ok(xs)
|
Ok(xs)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn forward_step(
|
pub fn forward_with_cache(
|
||||||
&mut self,
|
&mut self,
|
||||||
xs: &Tensor,
|
xs: &Tensor,
|
||||||
cos: &Tensor,
|
cos: &Tensor,
|
||||||
@@ -168,7 +168,7 @@ impl MiniCPMDecoderLayer {
|
|||||||
) -> Result<Tensor> {
|
) -> Result<Tensor> {
|
||||||
let residual = xs.clone();
|
let residual = xs.clone();
|
||||||
let xs = self.input_layernorm.forward(xs)?;
|
let xs = self.input_layernorm.forward(xs)?;
|
||||||
let xs = self.self_attn.forward_step(&xs, cos, sin, attention_mask, true)?;
|
let xs = self.self_attn.forward_with_cache(&xs, cos, sin, attention_mask, true)?;
|
||||||
let xs = (residual
|
let xs = (residual
|
||||||
+ xs.affine(
|
+ xs.affine(
|
||||||
self.scale_depth as f64 / (self.num_hidden_layers as f64).sqrt(),
|
self.scale_depth as f64 / (self.num_hidden_layers as f64).sqrt(),
|
||||||
@@ -254,7 +254,7 @@ impl MiniCPMModel {
|
|||||||
Ok(logits)
|
Ok(logits)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn forward_step(&mut self, input_ids: &Tensor, position_id: usize) -> Result<Tensor> {
|
pub fn forward_with_cache(&mut self, input_ids: &Tensor, position_id: usize) -> Result<Tensor> {
|
||||||
let (bs, seq_len) = input_ids.dims2()?;
|
let (bs, seq_len) = input_ids.dims2()?;
|
||||||
let input_embeds = self
|
let input_embeds = self
|
||||||
.embed_tokens
|
.embed_tokens
|
||||||
@@ -276,7 +276,7 @@ impl MiniCPMModel {
|
|||||||
let mut hidden_states = input_embeds;
|
let mut hidden_states = input_embeds;
|
||||||
for decode_layer in &mut self.layers {
|
for decode_layer in &mut self.layers {
|
||||||
hidden_states =
|
hidden_states =
|
||||||
decode_layer.forward_step(&hidden_states, &cos, &sin, attention_mask)?;
|
decode_layer.forward_with_cache(&hidden_states, &cos, &sin, attention_mask)?;
|
||||||
}
|
}
|
||||||
hidden_states = self.norm.forward(&hidden_states)?;
|
hidden_states = self.norm.forward(&hidden_states)?;
|
||||||
let hidden_state = hidden_states.narrow(1, seq_len - 1, 1)?;
|
let hidden_state = hidden_states.narrow(1, seq_len - 1, 1)?;
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
use anyhow::{Ok, Result};
|
use anyhow::{Ok, Result};
|
||||||
use candle_core::{D, Tensor};
|
use candle_core::{D, Tensor};
|
||||||
use candle_nn::{Conv1d, Conv1dConfig, ConvTranspose1d, ConvTranspose1dConfig, Module, VarBuilder};
|
use candle_nn::{Conv1d, Conv1dConfig, ConvTranspose1d, ConvTranspose1dConfig, Module, VarBuilder};
|
||||||
use std::{result::Result::Ok as StdOk, thread, time};
|
use std::{result::Result::Ok as StdOk};
|
||||||
|
|
||||||
pub struct CausalConv1d {
|
pub struct CausalConv1d {
|
||||||
conv1d: Conv1d,
|
conv1d: Conv1d,
|
||||||
@@ -40,7 +40,6 @@ pub struct CausalConvTranspose1d {
|
|||||||
conv_transpose1d: ConvTranspose1d,
|
conv_transpose1d: ConvTranspose1d,
|
||||||
padding: usize,
|
padding: usize,
|
||||||
output_padding: usize,
|
output_padding: usize,
|
||||||
config: ConvTranspose1dConfig,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
impl CausalConvTranspose1d {
|
impl CausalConvTranspose1d {
|
||||||
@@ -66,7 +65,6 @@ impl CausalConvTranspose1d {
|
|||||||
conv_transpose1d,
|
conv_transpose1d,
|
||||||
padding,
|
padding,
|
||||||
output_padding,
|
output_padding,
|
||||||
config
|
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
pub fn forward(&self, x: &Tensor) -> Result<Tensor> {
|
pub fn forward(&self, x: &Tensor) -> Result<Tensor> {
|
||||||
@@ -468,10 +466,10 @@ impl CausalDecoder {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub struct AudioVAE {
|
pub struct AudioVAE {
|
||||||
encoder_dim: usize,
|
// encoder_dim: usize,
|
||||||
encoder_rates: Vec<usize>,
|
// encoder_rates: Vec<usize>,
|
||||||
decoder_dim: usize,
|
// decoder_dim: usize,
|
||||||
decoder_rates: Vec<usize>,
|
// decoder_rates: Vec<usize>,
|
||||||
pub latent_dim: usize,
|
pub latent_dim: usize,
|
||||||
hop_length: usize,
|
hop_length: usize,
|
||||||
encoder: CausalEncoder,
|
encoder: CausalEncoder,
|
||||||
@@ -511,10 +509,10 @@ impl AudioVAE {
|
|||||||
)?;
|
)?;
|
||||||
let chunk_size = hop_length;
|
let chunk_size = hop_length;
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
encoder_dim,
|
// encoder_dim,
|
||||||
encoder_rates,
|
// encoder_rates,
|
||||||
decoder_dim,
|
// decoder_dim,
|
||||||
decoder_rates,
|
// decoder_rates,
|
||||||
latent_dim,
|
latent_dim,
|
||||||
hop_length,
|
hop_length,
|
||||||
encoder,
|
encoder,
|
||||||
|
|||||||
@@ -0,0 +1,161 @@
|
|||||||
|
use std::collections::HashMap;
|
||||||
|
|
||||||
|
use crate::{
|
||||||
|
models::voxcpm::{
|
||||||
|
audio_vae::AudioVAE, config::VoxCPMConfig, model::VoxCPMModel,
|
||||||
|
tokenizer::SingleChineseTokenizer,
|
||||||
|
},
|
||||||
|
utils::utils::{find_type_files, get_device, get_dtype},
|
||||||
|
};
|
||||||
|
use anyhow::{Ok, Result};
|
||||||
|
use candle_core::{DType, Device, Tensor, pickle::read_all_with_key};
|
||||||
|
use candle_nn::VarBuilder;
|
||||||
|
|
||||||
|
pub struct VoxCPMGenerate {
|
||||||
|
voxcpm: VoxCPMModel,
|
||||||
|
prompt_cache: Option<HashMap<String, Tensor>>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl VoxCPMGenerate {
|
||||||
|
pub fn init(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> {
|
||||||
|
let device = &get_device(device);
|
||||||
|
let config_path = path.to_string() + "/config.json";
|
||||||
|
let config: VoxCPMConfig = serde_json::from_slice(&std::fs::read(config_path)?)?;
|
||||||
|
let cfg_dtype = config.dtype.as_str();
|
||||||
|
let model_list = find_type_files(path, "pth")?;
|
||||||
|
println!(" pth model_list: {:?}", model_list);
|
||||||
|
let mut dict_to_hashmap = HashMap::new();
|
||||||
|
let mut vae_dtype = candle_core::DType::F32;
|
||||||
|
for m in model_list {
|
||||||
|
let dict = read_all_with_key(m, Some("state_dict"))?;
|
||||||
|
vae_dtype = dict[0].1.dtype();
|
||||||
|
for (k, v) in dict {
|
||||||
|
// println!("key: {}, tensor shape: {:?}", k, v);
|
||||||
|
dict_to_hashmap.insert(k, v);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
let vb_vae = VarBuilder::from_tensors(dict_to_hashmap, vae_dtype, &device);
|
||||||
|
let audio_vae = AudioVAE::new(
|
||||||
|
vb_vae,
|
||||||
|
128,
|
||||||
|
vec![2, 5, 8, 8],
|
||||||
|
Some(64),
|
||||||
|
1536,
|
||||||
|
vec![8, 8, 5, 2],
|
||||||
|
16000,
|
||||||
|
)?;
|
||||||
|
|
||||||
|
let model_list = find_type_files(path, "bin")?;
|
||||||
|
println!(" bin model_list: {:?}", model_list);
|
||||||
|
dict_to_hashmap = HashMap::new();
|
||||||
|
let mut m_dtype = get_dtype(dtype, cfg_dtype);
|
||||||
|
for m in model_list {
|
||||||
|
let dict = read_all_with_key(m, Some("state_dict"))?;
|
||||||
|
m_dtype = dict[0].1.dtype();
|
||||||
|
for (k, v) in dict {
|
||||||
|
// println!("key: {}, tensor shape: {:?}", k, v);
|
||||||
|
dict_to_hashmap.insert(k, v);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
let vb_voxcpm = VarBuilder::from_tensors(dict_to_hashmap, m_dtype, device);
|
||||||
|
let config_path = path.to_string() + "/config.json";
|
||||||
|
let config: VoxCPMConfig = serde_json::from_slice(&std::fs::read(config_path)?)?;
|
||||||
|
let tokenizer = SingleChineseTokenizer::new(path)?;
|
||||||
|
let voxcpm = VoxCPMModel::new(vb_voxcpm, config, tokenizer, audio_vae)?;
|
||||||
|
|
||||||
|
Ok(Self {
|
||||||
|
voxcpm,
|
||||||
|
prompt_cache: None,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn build_prompt_cache(
|
||||||
|
&mut self,
|
||||||
|
prompt_text: String,
|
||||||
|
prompt_wav_path: String,
|
||||||
|
) -> Result<()> {
|
||||||
|
let cache = self
|
||||||
|
.voxcpm
|
||||||
|
.build_prompt_cache(prompt_text, prompt_wav_path)?;
|
||||||
|
self.prompt_cache = Some(cache);
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn generate_use_prompt_cache(
|
||||||
|
&mut self,
|
||||||
|
target_text: String,
|
||||||
|
min_len: usize,
|
||||||
|
max_len: usize,
|
||||||
|
inference_timesteps: usize,
|
||||||
|
cfg_value: f64,
|
||||||
|
retry_badcase: bool,
|
||||||
|
retry_badcase_ratio_threshold: f64,
|
||||||
|
) -> Result<Tensor> {
|
||||||
|
let audio = match &self.prompt_cache {
|
||||||
|
Some(cache) => {
|
||||||
|
let prompt_cache = cache.clone();
|
||||||
|
self.voxcpm.generate_with_prompt_cache(
|
||||||
|
target_text,
|
||||||
|
prompt_cache,
|
||||||
|
min_len,
|
||||||
|
max_len,
|
||||||
|
inference_timesteps,
|
||||||
|
cfg_value,
|
||||||
|
retry_badcase,
|
||||||
|
retry_badcase_ratio_threshold,
|
||||||
|
)?
|
||||||
|
}
|
||||||
|
None => self.generate_simple(target_text)?,
|
||||||
|
};
|
||||||
|
Ok(audio)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn generate_with_prompt_simple(
|
||||||
|
&mut self,
|
||||||
|
target_text: String,
|
||||||
|
prompt_text: Option<String>,
|
||||||
|
prompt_wav_path: Option<String>,
|
||||||
|
) -> Result<Tensor> {
|
||||||
|
let audio = self.generate(
|
||||||
|
target_text,
|
||||||
|
prompt_text,
|
||||||
|
prompt_wav_path,
|
||||||
|
2,
|
||||||
|
1000,
|
||||||
|
10,
|
||||||
|
2.0,
|
||||||
|
false,
|
||||||
|
6.0,
|
||||||
|
)?;
|
||||||
|
Ok(audio)
|
||||||
|
}
|
||||||
|
pub fn generate_simple(&mut self, target_text: String) -> Result<Tensor> {
|
||||||
|
let audio = self.generate(target_text, None, None, 2, 1000, 10, 2.0, false, 6.0)?;
|
||||||
|
Ok(audio)
|
||||||
|
}
|
||||||
|
pub fn generate(
|
||||||
|
&mut self,
|
||||||
|
target_text: String,
|
||||||
|
prompt_text: Option<String>,
|
||||||
|
prompt_wav_path: Option<String>,
|
||||||
|
min_len: usize,
|
||||||
|
max_len: usize,
|
||||||
|
inference_timesteps: usize,
|
||||||
|
cfg_value: f64,
|
||||||
|
retry_badcase: bool,
|
||||||
|
retry_badcase_ratio_threshold: f64,
|
||||||
|
) -> Result<Tensor> {
|
||||||
|
let audio = self.voxcpm.generate(
|
||||||
|
target_text,
|
||||||
|
prompt_text,
|
||||||
|
prompt_wav_path,
|
||||||
|
min_len,
|
||||||
|
max_len,
|
||||||
|
inference_timesteps,
|
||||||
|
cfg_value,
|
||||||
|
retry_badcase,
|
||||||
|
retry_badcase_ratio_threshold,
|
||||||
|
)?;
|
||||||
|
Ok(audio)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,4 +1,3 @@
|
|||||||
use std::{thread, time};
|
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
@@ -10,7 +9,7 @@ use crate::{
|
|||||||
};
|
};
|
||||||
use anyhow::{anyhow, Ok, Result};
|
use anyhow::{anyhow, Ok, Result};
|
||||||
use candle_core::{DType, Device, Tensor, D};
|
use candle_core::{DType, Device, Tensor, D};
|
||||||
use candle_nn::{Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, rms_norm};
|
use candle_nn::{Embedding, Module, RmsNorm, VarBuilder, embedding, rms_norm};
|
||||||
|
|
||||||
pub struct MiniCPMLongRoPE {
|
pub struct MiniCPMLongRoPE {
|
||||||
short_factor: Vec<f32>,
|
short_factor: Vec<f32>,
|
||||||
@@ -177,7 +176,7 @@ impl MiniCPMDecoderLayer {
|
|||||||
Ok(xs)
|
Ok(xs)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn forward_step(
|
pub fn forward_with_cache(
|
||||||
&mut self,
|
&mut self,
|
||||||
xs: &Tensor,
|
xs: &Tensor,
|
||||||
cos: &Tensor,
|
cos: &Tensor,
|
||||||
@@ -188,7 +187,7 @@ impl MiniCPMDecoderLayer {
|
|||||||
let xs = self.input_layernorm.forward(xs)?;
|
let xs = self.input_layernorm.forward(xs)?;
|
||||||
let xs = self
|
let xs = self
|
||||||
.self_attn
|
.self_attn
|
||||||
.forward_step(&xs, cos, sin, attention_mask, true)?;
|
.forward_with_cache(&xs, cos, sin, attention_mask, true)?;
|
||||||
let xs = if self.use_mup {
|
let xs = if self.use_mup {
|
||||||
let res_add = (residual
|
let res_add = (residual
|
||||||
+ xs.affine(
|
+ xs.affine(
|
||||||
@@ -221,12 +220,11 @@ impl MiniCPMDecoderLayer {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub struct MiniCPMModel {
|
pub struct MiniCPMModel {
|
||||||
cfg: VoxMiniCPM4Config,
|
// cfg: VoxMiniCPM4Config,
|
||||||
pub embed_tokens: Option<Embedding>,
|
pub embed_tokens: Option<Embedding>,
|
||||||
layers: Vec<MiniCPMDecoderLayer>,
|
layers: Vec<MiniCPMDecoderLayer>,
|
||||||
norm: RmsNorm,
|
norm: RmsNorm,
|
||||||
rope_emb: MiniCPMLongRoPE,
|
rope_emb: MiniCPMLongRoPE,
|
||||||
// lm_head: Linear,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
impl MiniCPMModel {
|
impl MiniCPMModel {
|
||||||
@@ -250,23 +248,17 @@ impl MiniCPMModel {
|
|||||||
}
|
}
|
||||||
let norm = rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("norm"))?;
|
let norm = rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("norm"))?;
|
||||||
let rope_emb = MiniCPMLongRoPE::new(&cfg, vb.device(), vb.dtype())?;
|
let rope_emb = MiniCPMLongRoPE::new(&cfg, vb.device(), vb.dtype())?;
|
||||||
// let lm_head = Linear::new(embed_tokens.embeddings().clone(), None);
|
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
cfg,
|
// cfg,
|
||||||
embed_tokens,
|
embed_tokens,
|
||||||
layers,
|
layers,
|
||||||
norm,
|
norm,
|
||||||
rope_emb,
|
rope_emb,
|
||||||
// lm_head,
|
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn forward(&mut self, input_embeds: &Tensor, position_id: usize, is_causal: bool) -> Result<Tensor> {
|
pub fn forward(&mut self, input_embeds: &Tensor, position_id: usize, is_causal: bool) -> Result<Tensor> {
|
||||||
let (bs, seq_len, _) = input_embeds.dims3()?;
|
let (bs, seq_len, _) = input_embeds.dims3()?;
|
||||||
// let input_embeds = self
|
|
||||||
// .embed_tokens
|
|
||||||
// .forward(&input_ids)?
|
|
||||||
// .affine(self.cfg.scale_emb, 0.0)?;
|
|
||||||
let attention_mask: Option<&Tensor> = {
|
let attention_mask: Option<&Tensor> = {
|
||||||
if !is_causal || seq_len <= 1 {
|
if !is_causal || seq_len <= 1 {
|
||||||
None
|
None
|
||||||
@@ -288,17 +280,13 @@ impl MiniCPMModel {
|
|||||||
Ok(hidden_states)
|
Ok(hidden_states)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn forward_step(&mut self, input_embeds: &Tensor, position_id: usize) -> Result<Tensor> {
|
pub fn forward_with_cache(&mut self, input_embeds: &Tensor, position_id: usize) -> Result<Tensor> {
|
||||||
let input_embeds = match input_embeds.rank() {
|
let input_embeds = match input_embeds.rank() {
|
||||||
2 => input_embeds.unsqueeze(1)?,
|
2 => input_embeds.unsqueeze(1)?,
|
||||||
3 => input_embeds.clone(),
|
3 => input_embeds.clone(),
|
||||||
_ => return Err(anyhow!("MiniCPMModelinput_embeds illigal"))
|
_ => return Err(anyhow!("MiniCPMModelinput_embeds illigal"))
|
||||||
};
|
};
|
||||||
let (bs, seq_len, _) = input_embeds.dims3()?;
|
let (bs, seq_len, _) = input_embeds.dims3()?;
|
||||||
// let input_embeds = self
|
|
||||||
// .embed_tokens
|
|
||||||
// .forward(&input_ids)?
|
|
||||||
// .affine(self.cfg.scale_emb, 0.0)?;
|
|
||||||
let attention_mask: Option<&Tensor> = {
|
let attention_mask: Option<&Tensor> = {
|
||||||
if seq_len <= 1 {
|
if seq_len <= 1 {
|
||||||
None
|
None
|
||||||
@@ -315,7 +303,7 @@ impl MiniCPMModel {
|
|||||||
let mut hidden_states = input_embeds.clone();
|
let mut hidden_states = input_embeds.clone();
|
||||||
for decode_layer in &mut self.layers {
|
for decode_layer in &mut self.layers {
|
||||||
hidden_states =
|
hidden_states =
|
||||||
decode_layer.forward_step(&hidden_states, &cos, &sin, attention_mask)?;
|
decode_layer.forward_with_cache(&hidden_states, &cos, &sin, attention_mask)?;
|
||||||
}
|
}
|
||||||
hidden_states = self.norm.forward(&hidden_states)?;
|
hidden_states = self.norm.forward(&hidden_states)?;
|
||||||
|
|
||||||
|
|||||||
@@ -2,4 +2,5 @@ pub mod config;
|
|||||||
pub mod audio_vae;
|
pub mod audio_vae;
|
||||||
pub mod minicpm4;
|
pub mod minicpm4;
|
||||||
pub mod tokenizer;
|
pub mod tokenizer;
|
||||||
pub mod model;
|
pub mod model;
|
||||||
|
pub mod generate;
|
||||||
+158
-30
@@ -1,4 +1,4 @@
|
|||||||
use std::{cmp::max, f64, thread, time};
|
use std::{cmp::max, collections::HashMap, f64};
|
||||||
|
|
||||||
use anyhow::{Ok, Result};
|
use anyhow::{Ok, Result};
|
||||||
use candle_core::{D, DType, Device, IndexOp, Tensor};
|
use candle_core::{D, DType, Device, IndexOp, Tensor};
|
||||||
@@ -7,7 +7,7 @@ use candle_transformers::models::deepseek2::SplitOp;
|
|||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::voxcpm::{
|
models::voxcpm::{
|
||||||
audio_vae::{self, AudioVAE},
|
audio_vae::{AudioVAE},
|
||||||
config::{CfmConfig, VoxCPMConfig, VoxMiniCPM4Config},
|
config::{CfmConfig, VoxCPMConfig, VoxMiniCPM4Config},
|
||||||
minicpm4::MiniCPMModel,
|
minicpm4::MiniCPMModel,
|
||||||
tokenizer::SingleChineseTokenizer,
|
tokenizer::SingleChineseTokenizer,
|
||||||
@@ -118,8 +118,8 @@ pub struct VoxCPMLocDiT {
|
|||||||
time_mlp: TimestepEmbedding,
|
time_mlp: TimestepEmbedding,
|
||||||
delta_time_mlp: TimestepEmbedding,
|
delta_time_mlp: TimestepEmbedding,
|
||||||
decoder: MiniCPMModel,
|
decoder: MiniCPMModel,
|
||||||
config: VoxMiniCPM4Config,
|
// config: VoxMiniCPM4Config,
|
||||||
in_channels: usize,
|
// in_channels: usize,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl VoxCPMLocDiT {
|
impl VoxCPMLocDiT {
|
||||||
@@ -150,8 +150,8 @@ impl VoxCPMLocDiT {
|
|||||||
time_mlp,
|
time_mlp,
|
||||||
delta_time_mlp,
|
delta_time_mlp,
|
||||||
decoder,
|
decoder,
|
||||||
config,
|
// config,
|
||||||
in_channels,
|
// in_channels,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -188,9 +188,9 @@ impl VoxCPMLocDiT {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub struct UnifiedCFM {
|
pub struct UnifiedCFM {
|
||||||
solver: String,
|
// solver: String,
|
||||||
sigma_min: f32,
|
// sigma_min: f32,
|
||||||
t_scheduler: String,
|
// t_scheduler: String,
|
||||||
in_channels: usize,
|
in_channels: usize,
|
||||||
mean_mode: bool,
|
mean_mode: bool,
|
||||||
estimator: VoxCPMLocDiT,
|
estimator: VoxCPMLocDiT,
|
||||||
@@ -207,9 +207,9 @@ impl UnifiedCFM {
|
|||||||
let sigma_min = cfm_params.sigma_min;
|
let sigma_min = cfm_params.sigma_min;
|
||||||
let t_scheduler = cfm_params.t_scheduler;
|
let t_scheduler = cfm_params.t_scheduler;
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
solver,
|
// solver,
|
||||||
sigma_min,
|
// sigma_min,
|
||||||
t_scheduler,
|
// t_scheduler,
|
||||||
in_channels,
|
in_channels,
|
||||||
mean_mode,
|
mean_mode,
|
||||||
estimator,
|
estimator,
|
||||||
@@ -227,7 +227,7 @@ impl UnifiedCFM {
|
|||||||
sway_sampling_coef: f64,
|
sway_sampling_coef: f64,
|
||||||
use_cfg_zero_star: bool,
|
use_cfg_zero_star: bool,
|
||||||
) -> Result<Tensor> {
|
) -> Result<Tensor> {
|
||||||
let (b, c) = mu.dims2()?;
|
let (b, _) = mu.dims2()?;
|
||||||
let t = patch_size;
|
let t = patch_size;
|
||||||
let dtype = mu.dtype();
|
let dtype = mu.dtype();
|
||||||
let z = Tensor::randn(0.0f32, 1.0, (b, self.in_channels, t), mu.device())?
|
let z = Tensor::randn(0.0f32, 1.0, (b, self.in_channels, t), mu.device())?
|
||||||
@@ -345,7 +345,7 @@ impl VoxCPMLocEnc {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub fn forward(&mut self, x: &Tensor) -> Result<Tensor> {
|
pub fn forward(&mut self, x: &Tensor) -> Result<Tensor> {
|
||||||
let (b, t, p, d) = x.dims4()?;
|
let (b, t, _, _) = x.dims4()?;
|
||||||
let x = self.in_proj.forward(x)?;
|
let x = self.in_proj.forward(x)?;
|
||||||
let special_tokens = self.special_token.expand((b, t, 1, self.hidden_size))?;
|
let special_tokens = self.special_token.expand((b, t, 1, self.hidden_size))?;
|
||||||
let x = Tensor::cat(&[special_tokens, x], 2)?;
|
let x = Tensor::cat(&[special_tokens, x], 2)?;
|
||||||
@@ -362,7 +362,7 @@ pub struct VoxCPMModel {
|
|||||||
config: VoxCPMConfig,
|
config: VoxCPMConfig,
|
||||||
patch_size: usize,
|
patch_size: usize,
|
||||||
audio_start_token: usize,
|
audio_start_token: usize,
|
||||||
audio_end_token: usize,
|
// audio_end_token: usize,
|
||||||
chunk_size: usize,
|
chunk_size: usize,
|
||||||
sample_rate: usize,
|
sample_rate: usize,
|
||||||
tokenizer: SingleChineseTokenizer,
|
tokenizer: SingleChineseTokenizer,
|
||||||
@@ -390,7 +390,7 @@ impl VoxCPMModel {
|
|||||||
) -> Result<Self> {
|
) -> Result<Self> {
|
||||||
let base_lm = MiniCPMModel::new(vb.pp("base_lm"), config.lm_config.clone())?;
|
let base_lm = MiniCPMModel::new(vb.pp("base_lm"), config.lm_config.clone())?;
|
||||||
let audio_start_token = 101usize;
|
let audio_start_token = 101usize;
|
||||||
let audio_end_token = 102usize;
|
// let audio_end_token = 102usize;
|
||||||
let mut residual_lm_config = config.lm_config.clone();
|
let mut residual_lm_config = config.lm_config.clone();
|
||||||
residual_lm_config.num_hidden_layers = config.residual_lm_num_layers;
|
residual_lm_config.num_hidden_layers = config.residual_lm_num_layers;
|
||||||
residual_lm_config.vocab_size = 0;
|
residual_lm_config.vocab_size = 0;
|
||||||
@@ -456,7 +456,7 @@ impl VoxCPMModel {
|
|||||||
config,
|
config,
|
||||||
patch_size,
|
patch_size,
|
||||||
audio_start_token,
|
audio_start_token,
|
||||||
audio_end_token,
|
// audio_end_token,
|
||||||
chunk_size: audio_vae.chunk_size,
|
chunk_size: audio_vae.chunk_size,
|
||||||
sample_rate: audio_vae.sample_rate,
|
sample_rate: audio_vae.sample_rate,
|
||||||
tokenizer,
|
tokenizer,
|
||||||
@@ -486,7 +486,6 @@ impl VoxCPMModel {
|
|||||||
inference_timesteps: usize,
|
inference_timesteps: usize,
|
||||||
cfg_value: f64,
|
cfg_value: f64,
|
||||||
retry_badcase: bool,
|
retry_badcase: bool,
|
||||||
retry_badcase_max_times: usize,
|
|
||||||
retry_badcase_ratio_threshold: f64,
|
retry_badcase_ratio_threshold: f64,
|
||||||
) -> Result<Tensor> {
|
) -> Result<Tensor> {
|
||||||
let (text_token, text_mask, audio_feat, audio_mask) = match prompt_wav_path {
|
let (text_token, text_mask, audio_feat, audio_mask) = match prompt_wav_path {
|
||||||
@@ -554,17 +553,41 @@ impl VoxCPMModel {
|
|||||||
(text_token, text_mask, audio_feat, audio_mask)
|
(text_token, text_mask, audio_feat, audio_mask)
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
let text_token = text_token.unsqueeze(0)?;
|
|
||||||
let text_mask = text_mask.unsqueeze(0)?;
|
|
||||||
let audio_feat = audio_feat.unsqueeze(0)?.to_dtype(DType::BF16)?;
|
|
||||||
let audio_mask = audio_mask.unsqueeze(0)?;
|
|
||||||
let target_text_length = self.tokenizer.encode(target_text)?.len();
|
let target_text_length = self.tokenizer.encode(target_text)?.len();
|
||||||
let max_len = if retry_badcase {
|
let max_len = if retry_badcase {
|
||||||
(target_text_length as f64 * retry_badcase_ratio_threshold + 10.0) as usize
|
(target_text_length as f64 * retry_badcase_ratio_threshold + 10.0) as usize
|
||||||
} else {
|
} else {
|
||||||
max_len
|
max_len
|
||||||
};
|
};
|
||||||
|
let decode_audio = self._generate(
|
||||||
|
&text_token,
|
||||||
|
&text_mask,
|
||||||
|
&audio_feat,
|
||||||
|
&audio_mask,
|
||||||
|
min_len,
|
||||||
|
max_len,
|
||||||
|
inference_timesteps,
|
||||||
|
cfg_value,
|
||||||
|
)?;
|
||||||
|
Ok(decode_audio)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn _generate(
|
||||||
|
&mut self,
|
||||||
|
text_token: &Tensor,
|
||||||
|
text_mask: &Tensor,
|
||||||
|
audio_feat: &Tensor,
|
||||||
|
audio_mask: &Tensor,
|
||||||
|
min_len: usize,
|
||||||
|
max_len: usize,
|
||||||
|
inference_timesteps: usize,
|
||||||
|
cfg_value: f64,
|
||||||
|
) -> Result<Tensor> {
|
||||||
|
let text_token = text_token.unsqueeze(0)?;
|
||||||
|
let text_mask = text_mask.unsqueeze(0)?;
|
||||||
|
let audio_feat = audio_feat.unsqueeze(0)?.to_dtype(DType::BF16)?;
|
||||||
|
let audio_mask = audio_mask.unsqueeze(0)?;
|
||||||
|
|
||||||
let latent_pred = self.inference(
|
let latent_pred = self.inference(
|
||||||
&text_token,
|
&text_token,
|
||||||
&text_mask,
|
&text_mask,
|
||||||
@@ -584,7 +607,7 @@ impl VoxCPMModel {
|
|||||||
Ok(decode_audio)
|
Ok(decode_audio)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn inference(
|
fn inference(
|
||||||
&mut self,
|
&mut self,
|
||||||
text: &Tensor,
|
text: &Tensor,
|
||||||
text_mask: &Tensor,
|
text_mask: &Tensor,
|
||||||
@@ -595,7 +618,7 @@ impl VoxCPMModel {
|
|||||||
inference_timesteps: usize,
|
inference_timesteps: usize,
|
||||||
cfg_value: f64,
|
cfg_value: f64,
|
||||||
) -> Result<Tensor> {
|
) -> Result<Tensor> {
|
||||||
let (b, t, p, d) = feat.dims4()?;
|
let (_, t, _, _) = feat.dims4()?;
|
||||||
let feat_embed = self.feat_encoder.forward(feat)?; // [b, t, h_feat]
|
let feat_embed = self.feat_encoder.forward(feat)?; // [b, t, h_feat]
|
||||||
let feat_embed = self.enc_to_lm_proj.forward(&feat_embed)?;
|
let feat_embed = self.enc_to_lm_proj.forward(&feat_embed)?;
|
||||||
let scale_emb = if self.config.lm_config.use_mup {
|
let scale_emb = if self.config.lm_config.use_mup {
|
||||||
@@ -619,7 +642,7 @@ impl VoxCPMModel {
|
|||||||
let mut pred_feat_seq = Vec::new();
|
let mut pred_feat_seq = Vec::new();
|
||||||
let mut position_id = 0;
|
let mut position_id = 0;
|
||||||
let mut seq_len = t;
|
let mut seq_len = t;
|
||||||
let enc_outputs = self.base_lm.forward_step(&combined_embed, position_id)?;
|
let enc_outputs = self.base_lm.forward_with_cache(&combined_embed, position_id)?;
|
||||||
let enc_outputs = self
|
let enc_outputs = self
|
||||||
.fsq_layer
|
.fsq_layer
|
||||||
.forward(&enc_outputs)?
|
.forward(&enc_outputs)?
|
||||||
@@ -630,7 +653,7 @@ impl VoxCPMModel {
|
|||||||
|
|
||||||
let input_embeds =
|
let input_embeds =
|
||||||
enc_outputs.add(&feat_mask.unsqueeze(D::Minus1)?.broadcast_mul(&feat_embed)?)?;
|
enc_outputs.add(&feat_mask.unsqueeze(D::Minus1)?.broadcast_mul(&feat_embed)?)?;
|
||||||
let residual_enc_outputs = self.residual_lm.forward_step(&input_embeds, position_id)?;
|
let residual_enc_outputs = self.residual_lm.forward_with_cache(&input_embeds, position_id)?;
|
||||||
let mut residual_hidden = residual_enc_outputs.i((.., t - 1, ..))?;
|
let mut residual_hidden = residual_enc_outputs.i((.., t - 1, ..))?;
|
||||||
|
|
||||||
for i in 0..max_len {
|
for i in 0..max_len {
|
||||||
@@ -671,16 +694,16 @@ impl VoxCPMModel {
|
|||||||
seq_len = 1;
|
seq_len = 1;
|
||||||
lm_hidden = self
|
lm_hidden = self
|
||||||
.base_lm
|
.base_lm
|
||||||
.forward_step(&curr_embed.i((.., 0, ..))?, position_id)?
|
.forward_with_cache(&curr_embed.i((.., 0, ..))?, position_id)?
|
||||||
.squeeze(1)?;
|
.squeeze(1)?;
|
||||||
lm_hidden = self.fsq_layer.forward(&lm_hidden)?;
|
lm_hidden = self.fsq_layer.forward(&lm_hidden)?;
|
||||||
residual_hidden = self
|
residual_hidden = self
|
||||||
.residual_lm
|
.residual_lm
|
||||||
.forward_step(&lm_hidden.add(&curr_embed.i((.., 0, ..))?)?, position_id)?
|
.forward_with_cache(&lm_hidden.add(&curr_embed.i((.., 0, ..))?)?, position_id)?
|
||||||
.squeeze(1)?;
|
.squeeze(1)?;
|
||||||
}
|
}
|
||||||
let pred_seq = Tensor::cat(&pred_feat_seq, 1)?; // (b, t, p, d)
|
let pred_seq = Tensor::cat(&pred_feat_seq, 1)?; // (b, t, p, d)
|
||||||
let (b, t, p, d) = pred_seq.dims4()?;
|
let (b, _, _, d) = pred_seq.dims4()?;
|
||||||
let feat_pred = pred_seq
|
let feat_pred = pred_seq
|
||||||
.permute((0, 3, 1, 2))?
|
.permute((0, 3, 1, 2))?
|
||||||
.reshape((b, d, ()))?
|
.reshape((b, d, ()))?
|
||||||
@@ -689,4 +712,109 @@ impl VoxCPMModel {
|
|||||||
self.residual_lm.clear_kv_cache();
|
self.residual_lm.clear_kv_cache();
|
||||||
Ok(feat_pred)
|
Ok(feat_pred)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub fn build_prompt_cache(
|
||||||
|
&mut self,
|
||||||
|
prompt_text: String,
|
||||||
|
prompt_wav_path: String,
|
||||||
|
) -> Result<HashMap<String, Tensor>> {
|
||||||
|
let text_token = self.tokenizer.encode(prompt_text)?;
|
||||||
|
let text_token = Tensor::from_slice(&text_token, text_token.len(), &self.device)?;
|
||||||
|
let mut audio =
|
||||||
|
load_audio_with_resample(prompt_wav_path, self.device.clone(), Some(self.sample_rate))?;
|
||||||
|
let patch_len = self.patch_size * self.chunk_size;
|
||||||
|
if audio.dim(1)? % patch_len != 0 {
|
||||||
|
audio = audio.pad_with_zeros(D::Minus1, 0, patch_len - audio.dim(1)? % patch_len)?;
|
||||||
|
}
|
||||||
|
let audio_feat = self.audio_vae.encode(&audio, Some(self.sample_rate))?;
|
||||||
|
let audio_feat = audio_feat
|
||||||
|
.reshape((self.audio_vae.latent_dim, (), self.patch_size))?
|
||||||
|
.permute((1, 2, 0))?;
|
||||||
|
let dim0 = audio_feat.dim(0)? - 1;
|
||||||
|
let audio_feat = audio_feat.i(..dim0)?;
|
||||||
|
let mut hashmap = HashMap::new();
|
||||||
|
hashmap.insert("text_token".to_string(), text_token);
|
||||||
|
hashmap.insert("audio_feat".to_string(), audio_feat);
|
||||||
|
Ok(hashmap)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn generate_with_prompt_cache(
|
||||||
|
&mut self,
|
||||||
|
target_text: String,
|
||||||
|
prompt_cache: HashMap<String, Tensor>,
|
||||||
|
min_len: usize,
|
||||||
|
max_len: usize,
|
||||||
|
inference_timesteps: usize,
|
||||||
|
cfg_value: f64,
|
||||||
|
retry_badcase: bool,
|
||||||
|
retry_badcase_ratio_threshold: f64,
|
||||||
|
) -> Result<Tensor> {
|
||||||
|
let target_text_token = self.tokenizer.encode(target_text.clone())?;
|
||||||
|
let target_text_token =
|
||||||
|
Tensor::from_slice(&target_text_token, target_text_token.len(), &self.device)?;
|
||||||
|
let text_token = match prompt_cache.get("text_token") {
|
||||||
|
Some(token) => Tensor::cat(&[token, &target_text_token], 0)?,
|
||||||
|
None => target_text_token,
|
||||||
|
};
|
||||||
|
let audio_start = Tensor::new(vec![self.audio_start_token as u32], &self.device)?;
|
||||||
|
let text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?;
|
||||||
|
let text_length = text_token.dim(0)?;
|
||||||
|
let (audio_length, audio_feat) = match prompt_cache.get("audio_feat") {
|
||||||
|
Some(feat) => (feat.dim(0)?, Some(feat.clone())),
|
||||||
|
None => (0, None),
|
||||||
|
};
|
||||||
|
let (text_token, text_mask, audio_feat, audio_mask) = if audio_length > 0 {
|
||||||
|
let audio_feat = audio_feat.unwrap();
|
||||||
|
let audio_length = audio_feat.dim(0)?;
|
||||||
|
let text_pad_token = Tensor::zeros(audio_length, DType::U32, &self.device)?;
|
||||||
|
let text_token = Tensor::cat(&[text_token, text_pad_token], D::Minus1)?;
|
||||||
|
let audio_pad_feat = Tensor::zeros(
|
||||||
|
(text_length, self.patch_size, self.audio_vae.latent_dim),
|
||||||
|
audio_feat.dtype(),
|
||||||
|
&self.device,
|
||||||
|
)?;
|
||||||
|
let audio_feat = Tensor::cat(&[audio_pad_feat, audio_feat], 0)?;
|
||||||
|
let text_mask = Tensor::cat(
|
||||||
|
&[
|
||||||
|
Tensor::ones(text_length, self.dtype, &self.device)?,
|
||||||
|
Tensor::zeros(audio_length, self.dtype, &self.device)?,
|
||||||
|
],
|
||||||
|
D::Minus1,
|
||||||
|
)?;
|
||||||
|
let audio_mask = Tensor::cat(
|
||||||
|
&[
|
||||||
|
Tensor::zeros(text_length, self.dtype, &self.device)?,
|
||||||
|
Tensor::ones(audio_length, self.dtype, &self.device)?,
|
||||||
|
],
|
||||||
|
D::Minus1,
|
||||||
|
)?;
|
||||||
|
(text_token, text_mask, audio_feat, audio_mask)
|
||||||
|
} else {
|
||||||
|
let audio_feat = Tensor::zeros(
|
||||||
|
(text_length, self.patch_size, self.audio_vae.latent_dim),
|
||||||
|
DType::F32,
|
||||||
|
&self.device,
|
||||||
|
)?;
|
||||||
|
let text_mask = Tensor::ones(text_length, self.dtype, &self.device)?;
|
||||||
|
let audio_mask = Tensor::zeros(text_length, self.dtype, &self.device)?;
|
||||||
|
(text_token, text_mask, audio_feat, audio_mask)
|
||||||
|
};
|
||||||
|
let target_text_length = self.tokenizer.encode(target_text)?.len();
|
||||||
|
let max_len = if retry_badcase {
|
||||||
|
(target_text_length as f64 * retry_badcase_ratio_threshold + 10.0) as usize
|
||||||
|
} else {
|
||||||
|
max_len
|
||||||
|
};
|
||||||
|
let decode_audio = self._generate(
|
||||||
|
&text_token,
|
||||||
|
&text_mask,
|
||||||
|
&audio_feat,
|
||||||
|
&audio_mask,
|
||||||
|
min_len,
|
||||||
|
max_len,
|
||||||
|
inference_timesteps,
|
||||||
|
cfg_value,
|
||||||
|
)?;
|
||||||
|
Ok(decode_audio)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -226,7 +226,7 @@ pub fn load_audio<P: AsRef<Path>>(path: P, device: Device) -> Result<(Tensor, us
|
|||||||
let samples: Vec<f32> = match spec.sample_format {
|
let samples: Vec<f32> = match spec.sample_format {
|
||||||
SampleFormat::Int => {
|
SampleFormat::Int => {
|
||||||
// 将整数样本转换为浮点数 [-1.0, 1.0]
|
// 将整数样本转换为浮点数 [-1.0, 1.0]
|
||||||
println!("spec.bits_per_sample: {}", spec.bits_per_sample);
|
// println!("spec.bits_per_sample: {}", spec.bits_per_sample);
|
||||||
let samples = match spec.bits_per_sample {
|
let samples = match spec.bits_per_sample {
|
||||||
8 => {
|
8 => {
|
||||||
reader
|
reader
|
||||||
|
|||||||
@@ -223,7 +223,7 @@ pub fn masked_scatter_dim0(original: &Tensor, replace: &Tensor, mask: &Tensor) -
|
|||||||
let mask = mask.squeeze(0)?;
|
let mask = mask.squeeze(0)?;
|
||||||
let slices = nonzero_slice(&mask)?;
|
let slices = nonzero_slice(&mask)?;
|
||||||
let mut sub_start = 0usize;
|
let mut sub_start = 0usize;
|
||||||
let mut sub_end = 0usize;
|
let mut sub_end;
|
||||||
for (start, end) in slices {
|
for (start, end) in slices {
|
||||||
sub_end = sub_start + (end - start);
|
sub_end = sub_start + (end - start);
|
||||||
let sub_replace = replace.i((sub_start..sub_end, ..))?;
|
let sub_replace = replace.i((sub_start..sub_end, ..))?;
|
||||||
|
|||||||
@@ -21,7 +21,7 @@ fn minicpm_generate() -> Result<()> {
|
|||||||
"messages": [
|
"messages": [
|
||||||
{
|
{
|
||||||
"role": "user",
|
"role": "user",
|
||||||
"content": "贾宝玉和孙悟空有什么关系"
|
"content": "你好啊,你是谁"
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
}
|
}
|
||||||
|
|||||||
+34
-51
@@ -1,9 +1,9 @@
|
|||||||
use anyhow::{Ok, Result};
|
use anyhow::{Ok, Result};
|
||||||
use std::collections::HashMap;
|
use std::{collections::HashMap, time::Instant};
|
||||||
|
|
||||||
use aha::{
|
use aha::{
|
||||||
models::voxcpm::{
|
models::voxcpm::{
|
||||||
audio_vae::AudioVAE, config::VoxCPMConfig, model::VoxCPMModel,
|
audio_vae::AudioVAE, config::VoxCPMConfig, generate::VoxCPMGenerate, model::VoxCPMModel,
|
||||||
tokenizer::SingleChineseTokenizer,
|
tokenizer::SingleChineseTokenizer,
|
||||||
},
|
},
|
||||||
utils::{
|
utils::{
|
||||||
@@ -18,64 +18,47 @@ use candle_nn::VarBuilder;
|
|||||||
fn voxcpm_generate() -> Result<()> {
|
fn voxcpm_generate() -> Result<()> {
|
||||||
// RUST_BACKTRACE=1 cargo test -F cuda,flash-attn voxcpm_generate -- --nocapture
|
// RUST_BACKTRACE=1 cargo test -F cuda,flash-attn voxcpm_generate -- --nocapture
|
||||||
let model_path = "/home/jhq/huggingface_model/openbmb/VoxCPM-0.5B/";
|
let model_path = "/home/jhq/huggingface_model/openbmb/VoxCPM-0.5B/";
|
||||||
let model_list = find_type_files(&model_path, "pth")?;
|
|
||||||
println!(" pth model_list: {:?}", model_list);
|
let i_start = Instant::now();
|
||||||
let dev = get_device(None);
|
let mut voxcpm_generate = VoxCPMGenerate::init(model_path, None, None)?;
|
||||||
let mut dict_to_hashmap = HashMap::new();
|
let i_duration = i_start.elapsed();
|
||||||
let mut dtype = candle_core::DType::F32;
|
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||||
for m in model_list {
|
|
||||||
let dict = read_all_with_key(m, Some("state_dict"))?;
|
let i_start = Instant::now();
|
||||||
dtype = dict[0].1.dtype();
|
// let generate = voxcpm_generate.generate_simple("太阳当空照,花儿对我笑,小鸟说早早早".to_string())?;
|
||||||
for (k, v) in dict {
|
// let generate = voxcpm_generate.generate(
|
||||||
// println!("key: {}, tensor shape: {:?}", k, v);
|
// "太阳当空照,花儿对我笑,小鸟说早早早".to_string(),
|
||||||
// if k.contains("decoder.model.2.block.1") {
|
// Some("啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string()),
|
||||||
// println!("val: {}", v);
|
// Some("./assets/audio/voice_01.wav".to_string()),
|
||||||
// }
|
// // Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()),
|
||||||
dict_to_hashmap.insert(k, v);
|
// // Some("./assets/audio/voice_05.wav".to_string()),
|
||||||
}
|
// 2,
|
||||||
}
|
// 100,
|
||||||
let vb = VarBuilder::from_tensors(dict_to_hashmap, dtype, &dev);
|
// 10,
|
||||||
let audio_vae = AudioVAE::new(
|
// 2.0,
|
||||||
vb,
|
// false,
|
||||||
128,
|
// 6.0,
|
||||||
vec![2, 5, 8, 8],
|
// )?;
|
||||||
Some(64),
|
|
||||||
1536,
|
// 创建prompt_cache
|
||||||
vec![8, 8, 5, 2],
|
let _ = voxcpm_generate.build_prompt_cache(
|
||||||
16000,
|
"啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string(),
|
||||||
|
"./assets/audio/voice_01.wav".to_string(),
|
||||||
)?;
|
)?;
|
||||||
println!("audio vae load down");
|
// 使用prompt_cache生成语音
|
||||||
let model_list = find_type_files(&model_path, "bin")?;
|
let generate = voxcpm_generate.generate_use_prompt_cache(
|
||||||
println!(" bin model_list: {:?}", model_list);
|
|
||||||
dict_to_hashmap = HashMap::new();
|
|
||||||
for m in model_list {
|
|
||||||
let dict = read_all_with_key(m, Some("state_dict"))?;
|
|
||||||
dtype = dict[0].1.dtype();
|
|
||||||
for (k, v) in dict {
|
|
||||||
// println!("key: {}, tensor shape: {:?}", k, v);
|
|
||||||
dict_to_hashmap.insert(k, v);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
let vb_vox = VarBuilder::from_tensors(dict_to_hashmap, dtype, &dev);
|
|
||||||
let config_path = model_path.to_string() + "/config.json";
|
|
||||||
let config: VoxCPMConfig = serde_json::from_slice(&std::fs::read(config_path)?)?;
|
|
||||||
let tokenizer = SingleChineseTokenizer::new(model_path)?;
|
|
||||||
let mut voxcpm = VoxCPMModel::new(vb_vox, config, tokenizer, audio_vae)?;
|
|
||||||
let generate = voxcpm.generate(
|
|
||||||
"太阳当空照,花儿对我笑,小鸟说早早早".to_string(),
|
"太阳当空照,花儿对我笑,小鸟说早早早".to_string(),
|
||||||
Some("啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string()),
|
|
||||||
Some("./assets/audio/voice_01.wav".to_string()),
|
|
||||||
// Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()),
|
|
||||||
// Some("./assets/audio/voice_05.wav".to_string()),
|
|
||||||
2,
|
2,
|
||||||
100,
|
100,
|
||||||
10,
|
10,
|
||||||
2.0,
|
2.0,
|
||||||
false,
|
false,
|
||||||
3,
|
|
||||||
6.0,
|
6.0,
|
||||||
)?;
|
)?;
|
||||||
let _ = save_wav(&generate, "voxcpm_init.wav")?;
|
|
||||||
|
let i_duration = i_start.elapsed();
|
||||||
|
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||||
|
let _ = save_wav(&generate, "voxcpm.wav")?;
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
BIN
Binary file not shown.
Binary file not shown.
Reference in New Issue
Block a user