update voxcpm

This commit is contained in:
jhqxxx
2025-10-10 20:36:52 +08:00
parent fb746842d7
commit 39536df06b
14 changed files with 381 additions and 122 deletions
+1 -1
View File
@@ -179,7 +179,7 @@ impl AttentionNobias {
Ok(attn_output)
}
pub fn forward_step(
pub fn forward_with_cache(
&mut self,
xs: &Tensor,
cos: &Tensor,
+2 -2
View File
@@ -66,7 +66,7 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
None => 2048,
};
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 next_token = logit_processor.sample(&logits)?;
generate.push(next_token);
@@ -98,7 +98,7 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
let stream = stream! {
let mut error_tokens = Vec::new();
for _ in 0..sample_len {
let logits = self.minicpm.forward_step(
let logits = self.minicpm.forward_with_cache(
&input_ids,
seqlen_offset,
)?;
+4 -4
View File
@@ -159,7 +159,7 @@ impl MiniCPMDecoderLayer {
Ok(xs)
}
pub fn forward_step(
pub fn forward_with_cache(
&mut self,
xs: &Tensor,
cos: &Tensor,
@@ -168,7 +168,7 @@ impl MiniCPMDecoderLayer {
) -> Result<Tensor> {
let residual = xs.clone();
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
+ xs.affine(
self.scale_depth as f64 / (self.num_hidden_layers as f64).sqrt(),
@@ -254,7 +254,7 @@ impl MiniCPMModel {
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 input_embeds = self
.embed_tokens
@@ -276,7 +276,7 @@ impl MiniCPMModel {
let mut hidden_states = input_embeds;
for decode_layer in &mut self.layers {
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)?;
let hidden_state = hidden_states.narrow(1, seq_len - 1, 1)?;
+9 -11
View File
@@ -1,7 +1,7 @@
use anyhow::{Ok, Result};
use candle_core::{D, Tensor};
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 {
conv1d: Conv1d,
@@ -40,7 +40,6 @@ pub struct CausalConvTranspose1d {
conv_transpose1d: ConvTranspose1d,
padding: usize,
output_padding: usize,
config: ConvTranspose1dConfig,
}
impl CausalConvTranspose1d {
@@ -66,7 +65,6 @@ impl CausalConvTranspose1d {
conv_transpose1d,
padding,
output_padding,
config
})
}
pub fn forward(&self, x: &Tensor) -> Result<Tensor> {
@@ -468,10 +466,10 @@ impl CausalDecoder {
}
pub struct AudioVAE {
encoder_dim: usize,
encoder_rates: Vec<usize>,
decoder_dim: usize,
decoder_rates: Vec<usize>,
// encoder_dim: usize,
// encoder_rates: Vec<usize>,
// decoder_dim: usize,
// decoder_rates: Vec<usize>,
pub latent_dim: usize,
hop_length: usize,
encoder: CausalEncoder,
@@ -511,10 +509,10 @@ impl AudioVAE {
)?;
let chunk_size = hop_length;
Ok(Self {
encoder_dim,
encoder_rates,
decoder_dim,
decoder_rates,
// encoder_dim,
// encoder_rates,
// decoder_dim,
// decoder_rates,
latent_dim,
hop_length,
encoder,
+161
View File
@@ -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)
}
}
+7 -19
View File
@@ -1,4 +1,3 @@
use std::{thread, time};
use crate::{
models::{
@@ -10,7 +9,7 @@ use crate::{
};
use anyhow::{anyhow, Ok, Result};
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 {
short_factor: Vec<f32>,
@@ -177,7 +176,7 @@ impl MiniCPMDecoderLayer {
Ok(xs)
}
pub fn forward_step(
pub fn forward_with_cache(
&mut self,
xs: &Tensor,
cos: &Tensor,
@@ -188,7 +187,7 @@ impl MiniCPMDecoderLayer {
let xs = self.input_layernorm.forward(xs)?;
let xs = self
.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 res_add = (residual
+ xs.affine(
@@ -221,12 +220,11 @@ impl MiniCPMDecoderLayer {
}
pub struct MiniCPMModel {
cfg: VoxMiniCPM4Config,
// cfg: VoxMiniCPM4Config,
pub embed_tokens: Option<Embedding>,
layers: Vec<MiniCPMDecoderLayer>,
norm: RmsNorm,
rope_emb: MiniCPMLongRoPE,
// lm_head: Linear,
}
impl MiniCPMModel {
@@ -250,23 +248,17 @@ impl MiniCPMModel {
}
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 lm_head = Linear::new(embed_tokens.embeddings().clone(), None);
Ok(Self {
cfg,
// cfg,
embed_tokens,
layers,
norm,
rope_emb,
// lm_head,
})
}
pub fn forward(&mut self, input_embeds: &Tensor, position_id: usize, is_causal: bool) -> Result<Tensor> {
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> = {
if !is_causal || seq_len <= 1 {
None
@@ -288,17 +280,13 @@ impl MiniCPMModel {
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() {
2 => input_embeds.unsqueeze(1)?,
3 => input_embeds.clone(),
_ => return Err(anyhow!("MiniCPMModelinput_embeds illigal"))
};
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> = {
if seq_len <= 1 {
None
@@ -315,7 +303,7 @@ impl MiniCPMModel {
let mut hidden_states = input_embeds.clone();
for decode_layer in &mut self.layers {
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)?;
+2 -1
View File
@@ -2,4 +2,5 @@ pub mod config;
pub mod audio_vae;
pub mod minicpm4;
pub mod tokenizer;
pub mod model;
pub mod model;
pub mod generate;
+158 -30
View File
@@ -1,4 +1,4 @@
use std::{cmp::max, f64, thread, time};
use std::{cmp::max, collections::HashMap, f64};
use anyhow::{Ok, Result};
use candle_core::{D, DType, Device, IndexOp, Tensor};
@@ -7,7 +7,7 @@ use candle_transformers::models::deepseek2::SplitOp;
use crate::{
models::voxcpm::{
audio_vae::{self, AudioVAE},
audio_vae::{AudioVAE},
config::{CfmConfig, VoxCPMConfig, VoxMiniCPM4Config},
minicpm4::MiniCPMModel,
tokenizer::SingleChineseTokenizer,
@@ -118,8 +118,8 @@ pub struct VoxCPMLocDiT {
time_mlp: TimestepEmbedding,
delta_time_mlp: TimestepEmbedding,
decoder: MiniCPMModel,
config: VoxMiniCPM4Config,
in_channels: usize,
// config: VoxMiniCPM4Config,
// in_channels: usize,
}
impl VoxCPMLocDiT {
@@ -150,8 +150,8 @@ impl VoxCPMLocDiT {
time_mlp,
delta_time_mlp,
decoder,
config,
in_channels,
// config,
// in_channels,
})
}
@@ -188,9 +188,9 @@ impl VoxCPMLocDiT {
}
pub struct UnifiedCFM {
solver: String,
sigma_min: f32,
t_scheduler: String,
// solver: String,
// sigma_min: f32,
// t_scheduler: String,
in_channels: usize,
mean_mode: bool,
estimator: VoxCPMLocDiT,
@@ -207,9 +207,9 @@ impl UnifiedCFM {
let sigma_min = cfm_params.sigma_min;
let t_scheduler = cfm_params.t_scheduler;
Ok(Self {
solver,
sigma_min,
t_scheduler,
// solver,
// sigma_min,
// t_scheduler,
in_channels,
mean_mode,
estimator,
@@ -227,7 +227,7 @@ impl UnifiedCFM {
sway_sampling_coef: f64,
use_cfg_zero_star: bool,
) -> Result<Tensor> {
let (b, c) = mu.dims2()?;
let (b, _) = mu.dims2()?;
let t = patch_size;
let dtype = mu.dtype();
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> {
let (b, t, p, d) = x.dims4()?;
let (b, t, _, _) = x.dims4()?;
let x = self.in_proj.forward(x)?;
let special_tokens = self.special_token.expand((b, t, 1, self.hidden_size))?;
let x = Tensor::cat(&[special_tokens, x], 2)?;
@@ -362,7 +362,7 @@ pub struct VoxCPMModel {
config: VoxCPMConfig,
patch_size: usize,
audio_start_token: usize,
audio_end_token: usize,
// audio_end_token: usize,
chunk_size: usize,
sample_rate: usize,
tokenizer: SingleChineseTokenizer,
@@ -390,7 +390,7 @@ impl VoxCPMModel {
) -> Result<Self> {
let base_lm = MiniCPMModel::new(vb.pp("base_lm"), config.lm_config.clone())?;
let audio_start_token = 101usize;
let audio_end_token = 102usize;
// let audio_end_token = 102usize;
let mut residual_lm_config = config.lm_config.clone();
residual_lm_config.num_hidden_layers = config.residual_lm_num_layers;
residual_lm_config.vocab_size = 0;
@@ -456,7 +456,7 @@ impl VoxCPMModel {
config,
patch_size,
audio_start_token,
audio_end_token,
// audio_end_token,
chunk_size: audio_vae.chunk_size,
sample_rate: audio_vae.sample_rate,
tokenizer,
@@ -486,7 +486,6 @@ impl VoxCPMModel {
inference_timesteps: usize,
cfg_value: f64,
retry_badcase: bool,
retry_badcase_max_times: usize,
retry_badcase_ratio_threshold: f64,
) -> Result<Tensor> {
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)
}
};
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 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)
}
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(
&text_token,
&text_mask,
@@ -584,7 +607,7 @@ impl VoxCPMModel {
Ok(decode_audio)
}
pub fn inference(
fn inference(
&mut self,
text: &Tensor,
text_mask: &Tensor,
@@ -595,7 +618,7 @@ impl VoxCPMModel {
inference_timesteps: usize,
cfg_value: f64,
) -> 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.enc_to_lm_proj.forward(&feat_embed)?;
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 position_id = 0;
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
.fsq_layer
.forward(&enc_outputs)?
@@ -630,7 +653,7 @@ impl VoxCPMModel {
let input_embeds =
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, ..))?;
for i in 0..max_len {
@@ -671,16 +694,16 @@ impl VoxCPMModel {
seq_len = 1;
lm_hidden = self
.base_lm
.forward_step(&curr_embed.i((.., 0, ..))?, position_id)?
.forward_with_cache(&curr_embed.i((.., 0, ..))?, position_id)?
.squeeze(1)?;
lm_hidden = self.fsq_layer.forward(&lm_hidden)?;
residual_hidden = self
.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)?;
}
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
.permute((0, 3, 1, 2))?
.reshape((b, d, ()))?
@@ -689,4 +712,109 @@ impl VoxCPMModel {
self.residual_lm.clear_kv_cache();
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)
}
}
+1 -1
View File
@@ -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 {
SampleFormat::Int => {
// 将整数样本转换为浮点数 [-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 {
8 => {
reader
+1 -1
View File
@@ -223,7 +223,7 @@ pub fn masked_scatter_dim0(original: &Tensor, replace: &Tensor, mask: &Tensor) -
let mask = mask.squeeze(0)?;
let slices = nonzero_slice(&mask)?;
let mut sub_start = 0usize;
let mut sub_end = 0usize;
let mut sub_end;
for (start, end) in slices {
sub_end = sub_start + (end - start);
let sub_replace = replace.i((sub_start..sub_end, ..))?;
+1 -1
View File
@@ -21,7 +21,7 @@ fn minicpm_generate() -> Result<()> {
"messages": [
{
"role": "user",
"content": "贾宝玉和孙悟空有什么关系"
"content": "你好啊,你是谁"
}
]
}
+34 -51
View File
@@ -1,9 +1,9 @@
use anyhow::{Ok, Result};
use std::collections::HashMap;
use std::{collections::HashMap, time::Instant};
use aha::{
models::voxcpm::{
audio_vae::AudioVAE, config::VoxCPMConfig, model::VoxCPMModel,
audio_vae::AudioVAE, config::VoxCPMConfig, generate::VoxCPMGenerate, model::VoxCPMModel,
tokenizer::SingleChineseTokenizer,
},
utils::{
@@ -18,64 +18,47 @@ use candle_nn::VarBuilder;
fn voxcpm_generate() -> Result<()> {
// 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_list = find_type_files(&model_path, "pth")?;
println!(" pth model_list: {:?}", model_list);
let dev = get_device(None);
let mut dict_to_hashmap = HashMap::new();
let mut dtype = candle_core::DType::F32;
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);
// if k.contains("decoder.model.2.block.1") {
// println!("val: {}", v);
// }
dict_to_hashmap.insert(k, v);
}
}
let vb = VarBuilder::from_tensors(dict_to_hashmap, dtype, &dev);
let audio_vae = AudioVAE::new(
vb,
128,
vec![2, 5, 8, 8],
Some(64),
1536,
vec![8, 8, 5, 2],
16000,
let i_start = Instant::now();
let mut voxcpm_generate = VoxCPMGenerate::init(model_path, None, None)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
let i_start = Instant::now();
// let generate = voxcpm_generate.generate_simple("太阳当空照,花儿对我笑,小鸟说早早早".to_string())?;
// let generate = voxcpm_generate.generate(
// "太阳当空照,花儿对我笑,小鸟说早早早".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,
// 100,
// 10,
// 2.0,
// false,
// 6.0,
// )?;
// 创建prompt_cache
let _ = voxcpm_generate.build_prompt_cache(
"啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string(),
"./assets/audio/voice_01.wav".to_string(),
)?;
println!("audio vae load down");
let model_list = find_type_files(&model_path, "bin")?;
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(
// 使用prompt_cache生成语音
let generate = voxcpm_generate.generate_use_prompt_cache(
"太阳当空照,花儿对我笑,小鸟说早早早".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,
100,
10,
2.0,
false,
3,
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(())
}
BIN
View File
Binary file not shown.
BIN
View File
Binary file not shown.