refact voxcpm

This commit is contained in:
jhqxxx
2026-04-24 21:34:04 +08:00
parent 9e1adff2a1
commit 60a0ad69be
9 changed files with 962 additions and 26 deletions
+1
View File
@@ -23,6 +23,7 @@ pub mod qwen3_reranker;
pub mod qwen3vl; pub mod qwen3vl;
pub mod rmbg2_0; pub mod rmbg2_0;
pub mod voxcpm; pub mod voxcpm;
pub mod voxcpm_refact;
pub mod w2v_bert_2_0; pub mod w2v_bert_2_0;
// pub mod sam3; // pub mod sam3;
pub mod fire_red_vad; pub mod fire_red_vad;
+1
View File
@@ -265,6 +265,7 @@ impl GenerateModel for VoxCPMGenerate {
self.voxcpm.clear_kv_cache(); self.voxcpm.clear_kv_cache();
})?; })?;
let wav_u8 = get_audio_wav_u8(&audio, self.out_sample_rate as u32)?; let wav_u8 = get_audio_wav_u8(&audio, self.out_sample_rate as u32)?;
// let wave_u8_str = String::from_utf8(wav_u8)?;
let base64_audio = BASE64_STANDARD.encode(wav_u8); let base64_audio = BASE64_STANDARD.encode(wav_u8);
let response = build_audio_completion_response(&base64_audio, &self.model_name); let response = build_audio_completion_response(&base64_audio, &self.model_name);
self.voxcpm.clear_kv_cache(); self.voxcpm.clear_kv_cache();
+7 -1
View File
@@ -1,4 +1,5 @@
use anyhow::{Ok, Result, anyhow}; use anyhow::{Result, anyhow};
use candle_core::{Device, Tensor};
use tokenizers::Tokenizer; use tokenizers::Tokenizer;
pub struct SingleChineseTokenizer { pub struct SingleChineseTokenizer {
@@ -62,4 +63,9 @@ impl SingleChineseTokenizer {
.collect(); .collect();
Ok(ids) Ok(ids)
} }
pub fn encode_tensor(&self, text: String, device: &Device) -> Result<(Tensor, usize)> {
let ids = self.encode(text)?;
Ok((Tensor::from_slice(&ids, ids.len(), device)?, ids.len()))
}
} }
+225
View File
@@ -0,0 +1,225 @@
use anyhow::{Result, anyhow};
use candle_core::{DType, Device, Tensor, pickle::read_all_with_key};
use candle_nn::VarBuilder;
use rocket::futures::Stream;
use std::collections::HashMap;
use crate::{
models::{
voxcpm::{
audio_vae::AudioVAE,
config::{AudioVaeConfig, VoxCPMConfig},
tokenizer::SingleChineseTokenizer,
},
voxcpm_refact::{model::VoxCPMModelRefact, processor::VoxCPMProcessor},
},
utils::{find_type_files, get_device, get_dtype},
};
pub struct VoxCPMGenerateRefact {
voxcpm: VoxCPMModelRefact,
tokenizer: SingleChineseTokenizer,
audio_vae: AudioVAE,
processor: VoxCPMProcessor,
prompt_cache: Option<HashMap<String, Tensor>>,
out_sample_rate: usize,
// model_name: String,
}
impl VoxCPMGenerateRefact {
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 model_list = find_type_files(path, "pth")?;
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 {
dict_to_hashmap.insert(k, v);
}
}
let vb_vae = VarBuilder::from_tensors(dict_to_hashmap, vae_dtype, device);
let audio_config = match config.audio_vae_config.clone() {
Some(config) => config,
None => AudioVaeConfig {
encoder_dim: 128,
encoder_rates: vec![2, 5, 8, 8],
latent_dim: 64,
decoder_dim: 1536,
decoder_rates: vec![8, 8, 5, 2],
sample_rate: 16000,
out_sample_rate: None,
sr_bin_boundaries: None,
},
};
// let model_name = std::path::Path::new(path)
// .file_name()
// .and_then(|s| s.to_str())
// .unwrap_or("VoxCPM")
// .to_string();
let audio_vae = AudioVAE::new(
vb_vae,
audio_config.encoder_dim,
audio_config.encoder_rates.clone(),
Some(audio_config.latent_dim),
audio_config.decoder_dim,
audio_config.decoder_rates.clone(),
audio_config.sample_rate,
audio_config
.out_sample_rate
.unwrap_or(audio_config.sample_rate),
audio_config.sr_bin_boundaries,
Some("scale_bias".to_string()),
// Some(128),
// Some(false),
)?;
let processor = VoxCPMProcessor::new(
audio_vae.sample_rate,
audio_vae.chunk_size,
config.patch_size,
device.clone(),
);
let cfg_dtype = config.dtype.as_str();
let m_dtype = get_dtype(dtype, cfg_dtype);
let model_list = find_type_files(path, "bin")?;
// voxcpm0.5B模型文件是.bin类型, OpenBMB/VoxCPM1.5模型文件是.safetensors类型
let vb_voxcpm = if model_list.is_empty() {
let model_list = find_type_files(path, "safetensors")?;
unsafe { VarBuilder::from_mmaped_safetensors(&model_list, m_dtype, device)? }
} else {
dict_to_hashmap = HashMap::new();
let cfg_dtype = config.dtype.as_str();
let m_dtype = get_dtype(dtype, cfg_dtype);
for m in model_list {
let dict = read_all_with_key(m, Some("state_dict"))?;
for (k, v) in dict {
// println!("key: {}, tensor shape: {:?}", k, v);
dict_to_hashmap.insert(k, v);
}
}
VarBuilder::from_tensors(dict_to_hashmap, m_dtype, device)
};
let tokenizer = SingleChineseTokenizer::new(path)?;
let voxcpm = VoxCPMModelRefact::new(vb_voxcpm, config, audio_vae.latent_dim)?;
let out_sample_rate = audio_config
.out_sample_rate
.unwrap_or(audio_config.sample_rate);
Ok(Self {
voxcpm,
tokenizer,
audio_vae,
processor,
prompt_cache: None,
out_sample_rate,
// model_name,
})
}
pub fn sample_rate(&self) -> usize {
self.out_sample_rate
}
pub fn build_prompt_cache(
&mut self,
prompt_text: String,
prompt_wav_path: String,
) -> Result<()> {
let prompt_cache = self.processor.build_prompt_cache(
prompt_text,
prompt_wav_path,
&self.tokenizer,
&self.audio_vae,
)?;
self.prompt_cache = Some(prompt_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 (text_token, audio_feat, audio_mask) =
self.processor
.processor_use_cache(target_text, cache, &self.tokenizer)?;
let target_text_length = if let Some(mask) = &audio_mask {
text_token.dim(1)? - (mask.sum_all()?.to_scalar::<u32>()? as usize)
} else {
text_token.dim(1)?
};
let max_len = if retry_badcase {
(target_text_length as f64 * retry_badcase_ratio_threshold + 10.0) as usize
} else {
max_len
};
self.voxcpm.inference(
&text_token,
audio_feat.as_ref(),
audio_mask.as_ref(),
min_len,
max_len,
inference_timesteps,
cfg_value,
&self.audio_vae,
)?
}
None => {
return Err(anyhow!("need prompt_cache"));
}
};
self.voxcpm.clear_kv_cache();
Ok(audio)
}
pub fn generate_stream_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<impl Stream<Item = Result<Tensor, anyhow::Error>>> {
match &self.prompt_cache {
Some(cache) => {
let (text_token, audio_feat, audio_mask) =
self.processor
.processor_use_cache(target_text, cache, &self.tokenizer)?;
let target_text_length = if let Some(mask) = &audio_mask {
text_token.dim(1)? - (mask.sum_all()?.to_scalar::<u32>()? as usize)
} else {
text_token.dim(1)?
};
let max_len = if retry_badcase {
(target_text_length as f64 * retry_badcase_ratio_threshold + 10.0) as usize
} else {
max_len
};
self.voxcpm.inference_stream(
text_token,
audio_feat,
audio_mask,
min_len,
max_len,
inference_timesteps,
cfg_value,
&self.audio_vae,
)
}
None => Err(anyhow!("need prompt_cache")),
}
}
}
+3
View File
@@ -0,0 +1,3 @@
pub mod generate;
pub mod model;
pub mod processor;
+469
View File
@@ -0,0 +1,469 @@
use crate::{
models::voxcpm::{
audio_vae::AudioVAE,
config::VoxCPMConfig,
minicpm4::MiniCPMModel,
model::{ScalarQuantizationLayer, UnifiedCFM, VoxCPMLocDiT, VoxCPMLocEnc},
},
utils::tensor_utils::masked_scatter_dim0,
};
use anyhow::Result;
use candle_core::{D, DType, Device, IndexOp, Tensor};
use candle_nn::{Linear, Module, VarBuilder, linear, linear_no_bias};
use rocket::async_stream::stream;
use rocket::futures::Stream;
pub struct VoxCPMModelRefact {
config: VoxCPMConfig,
patch_size: usize,
latent_dim: usize,
// audio_start_token: u32,
// // audio_end_token: u32,
// ref_audio_start_token: u32,
// ref_audio_end_token: u32,
// chunk_size: usize,
// sample_rate: usize,
base_lm: MiniCPMModel,
residual_lm: MiniCPMModel,
feat_encoder: VoxCPMLocEnc,
feat_decoder: UnifiedCFM,
fsq_layer: ScalarQuantizationLayer,
enc_to_lm_proj: Linear,
lm_to_dit_proj: Linear,
res_to_dit_proj: Linear,
fusion_concat_proj: Option<Linear>,
stop_proj: Linear,
stop_head: Linear,
device: Device,
dtype: DType,
}
impl VoxCPMModelRefact {
pub fn new(vb: VarBuilder, config: VoxCPMConfig, latent_dim: usize) -> Result<Self> {
let base_lm = MiniCPMModel::new(vb.pp("base_lm"), 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.vocab_size = 0;
residual_lm_config.no_rope = config.residual_lm_no_rope;
let residual_lm = MiniCPMModel::new(vb.pp("residual_lm"), residual_lm_config)?;
let mut encoder_config = config.lm_config.clone();
encoder_config.hidden_size = config.encoder_config.hidden_dim;
encoder_config.intermediate_size = config.encoder_config.ffn_dim;
encoder_config.num_attention_heads = config.encoder_config.num_heads;
encoder_config.num_hidden_layers = config.encoder_config.num_layers;
encoder_config.kv_channels = config.encoder_config.kv_channels;
encoder_config.vocab_size = 0;
let feat_encoder =
VoxCPMLocEnc::new(vb.pp("feat_encoder"), encoder_config, config.feat_dim)?;
let mut decoder_config = config.lm_config.clone();
decoder_config.hidden_size = config.dit_config.hidden_dim;
decoder_config.intermediate_size = config.dit_config.ffn_dim;
decoder_config.num_attention_heads = config.dit_config.num_heads;
decoder_config.num_hidden_layers = config.dit_config.num_layers;
decoder_config.kv_channels = config.dit_config.kv_channels;
decoder_config.vocab_size = 0;
let estimator = VoxCPMLocDiT::new(
vb.pp("feat_decoder.estimator"),
decoder_config,
config.feat_dim,
)?;
let feat_decoder = UnifiedCFM::new(
config.feat_dim,
config.dit_config.cfm_config.clone(),
estimator,
false,
// config.architecture.clone(),
)?;
let fsq_layer = ScalarQuantizationLayer::new(
vb.pp("fsq_layer"),
config.lm_config.hidden_size,
config.lm_config.hidden_size,
config.scalar_quantization_latent_dim,
config.scalar_quantization_scale,
)?;
let enc_to_lm_proj = linear(
config.encoder_config.hidden_dim,
config.lm_config.hidden_size,
vb.pp("enc_to_lm_proj"),
)?;
let lm_to_dit_proj = linear(
config.lm_config.hidden_size,
config.dit_config.hidden_dim,
vb.pp("lm_to_dit_proj"),
)?;
let res_to_dit_proj = linear(
config.lm_config.hidden_size,
config.dit_config.hidden_dim,
vb.pp("res_to_dit_proj"),
)?;
let fusion_concat_proj = if config.architecture.to_lowercase().eq("voxcpm2") {
Some(linear(
config.lm_config.hidden_size * 2,
config.lm_config.hidden_size,
vb.pp("fusion_concat_proj"),
)?)
} else {
None
};
let stop_proj = linear(
config.lm_config.hidden_size,
config.lm_config.hidden_size,
vb.pp("stop_proj"),
)?;
let stop_head = linear_no_bias(config.lm_config.hidden_size, 2, vb.pp("stop_head"))?;
let patch_size = config.patch_size;
Ok(Self {
config,
patch_size,
latent_dim,
// audio_start_token: 101,
// // audio_end_token: 102,
// ref_audio_start_token: 103,
// ref_audio_end_token: 104,
// chunk_size: audio_vae.chunk_size,
// sample_rate: audio_vae.sample_rate,
// tokenizer,
// audio_vae,
base_lm,
residual_lm,
feat_encoder,
feat_decoder,
fsq_layer,
enc_to_lm_proj,
lm_to_dit_proj,
res_to_dit_proj,
fusion_concat_proj,
stop_proj,
stop_head,
device: vb.device().clone(),
dtype: vb.dtype(),
})
}
pub fn inference(
&mut self,
text: &Tensor,
audio_feat: Option<&Tensor>,
audio_mask: Option<&Tensor>,
min_len: usize,
max_len: usize,
inference_timesteps: usize,
cfg_value: f64,
audio_vae: &AudioVAE,
) -> Result<Tensor> {
let (b, t) = text.dims2()?;
let scale_emb = if self.config.lm_config.use_mup {
self.config.lm_config.scale_emb
} else {
1.0
};
let text_embed = self
.base_lm
.embed_tokens
.as_ref()
.unwrap()
.forward(text)?
.affine(scale_emb as f64, 0.0)?;
let (combined_embed, mut prefix_feat_cond, feat_embed) = if let Some(audio_feat) =
audio_feat
&& let Some(audio_mask) = audio_mask
{
let audio_feat = audio_feat.to_dtype(self.dtype)?;
let audio_t = audio_feat.dim(1)?;
let feat_embed = self.feat_encoder.forward(&audio_feat)?; // [b, audio_t, h_feat]
let feat_embed = self.enc_to_lm_proj.forward(&feat_embed)?.squeeze(0)?;
let embeds = masked_scatter_dim0(&text_embed, &feat_embed, audio_mask)?;
let prefix_feat_cond = audio_feat.i((.., audio_t - 1, ..))?;
(embeds, prefix_feat_cond, Some(feat_embed))
} else {
let prefix_feat_cond = Tensor::zeros(
(b, self.patch_size, self.latent_dim),
self.dtype,
&self.device,
)?;
(text_embed, prefix_feat_cond, None)
};
let mut pred_feat_seq = Vec::new();
// if feat_mask.i((1, t-1))?.to_scalar::<f32>()? == 0.0 {
// // TODO for stream
// }
let mut position_id = 0;
let mut seq_len = t;
let enc_outputs = self
.base_lm
.forward_with_cache(&combined_embed, position_id)?;
let (mut lm_hidden, input_embeds) = if let Some(_) = audio_feat
&& let Some(audio_mask) = audio_mask
&& let Some(feat_embed) = feat_embed
{
let fsq_emb = self.fsq_layer.forward(&enc_outputs)?;
let audio_mask_broadcast = audio_mask
.unsqueeze(D::Minus1)?
.broadcast_as(fsq_emb.shape())?;
let enc_outputs = audio_mask_broadcast.where_cond(&fsq_emb, &enc_outputs)?;
let lm_hidden = enc_outputs.i((.., t - 1, ..))?;
let input_embeds = if let Some(fusion) = &self.fusion_concat_proj {
let feat = enc_outputs.zeros_like()?;
let feat = masked_scatter_dim0(&feat, &feat_embed, audio_mask)?;
let concat = Tensor::cat(&[&enc_outputs, &feat], D::Minus1)?;
fusion.forward(&concat)?
} else {
let feat = enc_outputs.zeros_like()?;
let feat = masked_scatter_dim0(&feat, &feat_embed, audio_mask)?;
enc_outputs.add(&feat)?
};
(lm_hidden, input_embeds)
} else {
let lm_hidden = enc_outputs.i((.., t - 1, ..))?;
let input_embeds = if let Some(fusion) = &self.fusion_concat_proj {
let feat = enc_outputs.zeros_like()?;
let concat = Tensor::cat(&[&enc_outputs, &feat], D::Minus1)?;
fusion.forward(&concat)?
} else {
enc_outputs
};
(lm_hidden, input_embeds)
};
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 {
let dit_hidden_1 = self.lm_to_dit_proj.forward(&lm_hidden)?; // [b, h_dit]
let dit_hidden_2 = self.res_to_dit_proj.forward(&residual_hidden)?; // [b, h_dit]
// let dit_hidden = dit_hidden_1.add(&dit_hidden_2)?;
let dit_hidden = if self.fusion_concat_proj.is_some() {
Tensor::cat(&[&dit_hidden_1, &dit_hidden_2], D::Minus1)?
} else {
dit_hidden_1.add(&dit_hidden_2)?
};
let cond = prefix_feat_cond.transpose(1, 2)?.contiguous()?;
let pred_feat = self
.feat_decoder
.forward(
&dit_hidden,
inference_timesteps,
self.patch_size,
&cond,
1.0,
cfg_value,
1.0,
true,
)?
.transpose(1, 2)?; // [b, p, d]
let curr_embed = self.feat_encoder.forward(&pred_feat.unsqueeze(1)?)?; // [b, 1, c]
let curr_embed = self.enc_to_lm_proj.forward(&curr_embed)?;
pred_feat_seq.push(pred_feat.unsqueeze(1)?);
prefix_feat_cond = pred_feat;
let stop_flag = self.stop_proj.forward(&lm_hidden)?.silu()?;
let stop_flag = self
.stop_head
.forward(&stop_flag)?
.argmax(D::Minus1)?
.i(0)?
.to_scalar::<u32>()?;
if i > min_len && stop_flag == 1 {
break;
}
position_id += seq_len;
seq_len = 1;
lm_hidden = self
.base_lm
.forward_with_cache(&curr_embed.i((.., 0, ..))?, position_id)?
.squeeze(1)?;
lm_hidden = self.fsq_layer.forward(&lm_hidden)?;
let curr_residual_input = if let Some(fusion) = &self.fusion_concat_proj {
let curr_embed = curr_embed.i((.., 0, ..))?;
let concat = Tensor::cat(&[&lm_hidden, &curr_embed], D::Minus1)?;
fusion.forward(&concat)?
} else {
lm_hidden.add(&curr_embed.i((.., 0, ..))?)?
};
residual_hidden = self
.residual_lm
.forward_with_cache(&curr_residual_input, position_id)?
.squeeze(1)?;
}
let pred_seq = Tensor::cat(&pred_feat_seq, 1)?; // (b, t, p, d)
let (b, _, _, d) = pred_seq.dims4()?;
let feat_pred = pred_seq
.permute((0, 3, 1, 2))?
.reshape((b, d, ()))?
.contiguous()?;
self.clear_kv_cache();
let decode_audio = audio_vae
.decode(&feat_pred.to_dtype(DType::F32)?, None)?
.squeeze(1)?;
let decode_audio_len = decode_audio.dim(D::Minus1)? - 640 - 640;
let decode_audio = decode_audio.narrow(D::Minus1, 640, decode_audio_len)?;
Ok(decode_audio)
}
pub fn inference_stream(
&mut self,
text: Tensor,
audio_feat: Option<Tensor>,
audio_mask: Option<Tensor>,
min_len: usize,
max_len: usize,
inference_timesteps: usize,
cfg_value: f64,
audio_vae: &AudioVAE,
) -> Result<impl Stream<Item = Result<Tensor, anyhow::Error>>> {
let (b, t) = text.dims2()?;
let scale_emb = if self.config.lm_config.use_mup {
self.config.lm_config.scale_emb
} else {
1.0
};
let text_embed = self
.base_lm
.embed_tokens
.as_ref()
.unwrap()
.forward(&text)?
.affine(scale_emb as f64, 0.0)?;
let (combined_embed, mut prefix_feat_cond, feat_embed) = if let Some(audio_feat) =
&audio_feat
&& let Some(audio_mask) = &audio_mask
{
let audio_feat = audio_feat.to_dtype(self.dtype)?;
let audio_t = audio_feat.dim(1)?;
let feat_embed = self.feat_encoder.forward(&audio_feat)?; // [b, audio_t, h_feat]
let feat_embed = self.enc_to_lm_proj.forward(&feat_embed)?.squeeze(0)?;
let embeds = masked_scatter_dim0(&text_embed, &feat_embed, audio_mask)?;
let prefix_feat_cond = audio_feat.i((.., audio_t - 1, ..))?;
(embeds, prefix_feat_cond, Some(feat_embed))
} else {
let prefix_feat_cond = Tensor::zeros(
(b, self.patch_size, self.latent_dim),
self.dtype,
&self.device,
)?;
(text_embed, prefix_feat_cond, None)
};
// let mut pred_feat_seq = Vec::new();
// if feat_mask.i((1, t-1))?.to_scalar::<f32>()? == 0.0 {
// // TODO for stream
// }
let mut position_id = 0;
let mut seq_len = t;
let enc_outputs = self
.base_lm
.forward_with_cache(&combined_embed, position_id)?;
let (mut lm_hidden, input_embeds) = if let Some(_) = &audio_feat
&& let Some(audio_mask) = &audio_mask
&& let Some(feat_embed) = feat_embed
{
let fsq_emb = self.fsq_layer.forward(&enc_outputs)?;
let audio_mask_broadcast = audio_mask
.unsqueeze(D::Minus1)?
.broadcast_as(fsq_emb.shape())?;
let enc_outputs = audio_mask_broadcast.where_cond(&fsq_emb, &enc_outputs)?;
let lm_hidden = enc_outputs.i((.., t - 1, ..))?;
let input_embeds = if let Some(fusion) = &self.fusion_concat_proj {
let feat = enc_outputs.zeros_like()?;
let feat = masked_scatter_dim0(&feat, &feat_embed, audio_mask)?;
let concat = Tensor::cat(&[&enc_outputs, &feat], D::Minus1)?;
fusion.forward(&concat)?
} else {
let feat = enc_outputs.zeros_like()?;
let feat = masked_scatter_dim0(&feat, &feat_embed, audio_mask)?;
enc_outputs.add(&feat)?
};
(lm_hidden, input_embeds)
} else {
let lm_hidden = enc_outputs.i((.., t - 1, ..))?;
let input_embeds = if let Some(fusion) = &self.fusion_concat_proj {
let feat = enc_outputs.zeros_like()?;
let concat = Tensor::cat(&[&enc_outputs, &feat], D::Minus1)?;
fusion.forward(&concat)?
} else {
enc_outputs
};
(lm_hidden, input_embeds)
};
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 stream = stream! {
for i in 0..max_len {
let dit_hidden_1 = self.lm_to_dit_proj.forward(&lm_hidden)?; // [b, h_dit]
let dit_hidden_2 = self.res_to_dit_proj.forward(&residual_hidden)?; // [b, h_dit]
// let dit_hidden = dit_hidden_1.add(&dit_hidden_2)?;
let dit_hidden = if self.fusion_concat_proj.is_some() {
Tensor::cat(&[&dit_hidden_1, &dit_hidden_2], D::Minus1)?
} else {
dit_hidden_1.add(&dit_hidden_2)?
};
let cond = prefix_feat_cond.transpose(1, 2)?.contiguous()?;
let pred_feat = self
.feat_decoder
.forward(
&dit_hidden,
inference_timesteps,
self.patch_size,
&cond,
1.0,
cfg_value,
1.0,
true,
)?
.transpose(1, 2)?; // [b, p, d]
let curr_embed = self.feat_encoder.forward(&pred_feat.unsqueeze(1)?)?; // [b, 1, c]
let curr_embed = self.enc_to_lm_proj.forward(&curr_embed)?;
let single_feat_pred = pred_feat.permute((0, 2, 1))?.contiguous()?;
let decode_audio = audio_vae
.decode(&single_feat_pred.to_dtype(DType::F32)?, None)?
.squeeze(1)?;
yield Ok(decode_audio);
prefix_feat_cond = pred_feat;
let stop_flag = self.stop_proj.forward(&lm_hidden)?.silu()?;
let stop_flag = self
.stop_head
.forward(&stop_flag)?
.argmax(D::Minus1)?
.i(0)?
.to_scalar::<u32>()?;
if i > min_len && stop_flag == 1 {
break;
}
position_id += seq_len;
seq_len = 1;
lm_hidden = self
.base_lm
.forward_with_cache(&curr_embed.i((.., 0, ..))?, position_id)?
.squeeze(1)?;
lm_hidden = self.fsq_layer.forward(&lm_hidden)?;
let curr_residual_input = if let Some(fusion) = &self.fusion_concat_proj {
let curr_embed = curr_embed.i((.., 0, ..))?;
let concat = Tensor::cat(&[&lm_hidden, &curr_embed], D::Minus1)?;
fusion.forward(&concat)?
} else {
lm_hidden.add(&curr_embed.i((.., 0, ..))?)?
};
residual_hidden = self
.residual_lm
.forward_with_cache(&curr_residual_input, position_id)?
.squeeze(1)?;
}
self.clear_kv_cache();
};
Ok(stream)
}
pub fn clear_kv_cache(&mut self) {
self.base_lm.clear_kv_cache();
self.residual_lm.clear_kv_cache();
}
}
+165
View File
@@ -0,0 +1,165 @@
use std::collections::HashMap;
use crate::{
models::voxcpm::{audio_vae::AudioVAE, tokenizer::SingleChineseTokenizer},
utils::audio_utils::load_audio_with_resample,
};
use anyhow::Result;
use candle_core::{D, Device, IndexOp, Tensor};
pub struct VoxCPMProcessor {
sample_rate: usize,
chunk_size: usize,
patch_size: usize,
audio_start_token: u32,
ref_audio_start_token: u32,
ref_audio_end_token: u32,
device: Device,
}
impl VoxCPMProcessor {
pub fn new(sample_rate: usize, chunk_size: usize, patch_size: usize, device: Device) -> Self {
Self {
sample_rate,
chunk_size,
patch_size,
audio_start_token: 101,
ref_audio_start_token: 103,
ref_audio_end_token: 104,
device,
}
}
pub fn build_prompt_cache(
&mut self,
prompt_text: String,
prompt_wav_path: String,
tokenizer: &SingleChineseTokenizer,
audio_vae: &AudioVAE,
) -> Result<HashMap<String, Tensor>> {
let (text_token, _) = tokenizer.encode_tensor(prompt_text, &self.device)?;
let mut audio =
load_audio_with_resample(&prompt_wav_path, &self.device, 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 = audio_vae.encode(&audio, Some(self.sample_rate))?;
let audio_feat = audio_feat
.reshape((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 processor(
&self,
target_text: String,
prompt_text: Option<String>,
prompt_wav_path: Option<String>,
tokenizer: &SingleChineseTokenizer,
audio_vae: &AudioVAE,
) -> Result<(Tensor, Option<Tensor>, Option<Tensor>)> {
let text = if let Some(prompt_text) = &prompt_text {
prompt_text.clone() + &target_text
} else {
target_text
};
let (text_token, _) = tokenizer.encode_tensor(text, &self.device)?;
let audio_start = Tensor::new(vec![self.audio_start_token], &self.device)?;
let mut text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?;
let (audio_feat, audio_mask) = if let Some(path) = prompt_wav_path {
let mut audio = load_audio_with_resample(&path, &self.device, 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, patch_len - audio.dim(1)? % patch_len, 0)?;
}
let audio_feat = audio_vae.encode(&audio, Some(self.sample_rate))?;
let audio_feat = audio_feat
.reshape((audio_vae.latent_dim, (), self.patch_size))?
.permute((1, 2, 0))?;
let text_length = text_token.dim(0)?;
let audio_length = audio_feat.dim(0)?;
let audio_mask = if prompt_text.is_some() {
let text_pad_token =
Tensor::zeros(audio_length, candle_core::DType::U32, &self.device)?;
text_token = Tensor::cat(&[text_token, text_pad_token], D::Minus1)?;
let mask = Tensor::cat(
&[
Tensor::zeros(text_length, candle_core::DType::U32, &self.device)?,
Tensor::ones(audio_length, candle_core::DType::U32, &self.device)?,
],
D::Minus1,
)?
.unsqueeze(0)?;
Some(mask)
} else {
let ref_start = Tensor::new(vec![self.ref_audio_start_token], &self.device)?;
let ref_end = Tensor::new(vec![self.ref_audio_end_token], &self.device)?;
let ref_token = Tensor::zeros(audio_length, candle_core::DType::U32, &self.device)?;
text_token = Tensor::cat(&[&ref_start, &ref_token, &ref_end, &text_token], 0)?;
let mask = Tensor::cat(
&[
Tensor::new(vec![0u32], &self.device)?,
Tensor::ones(audio_length, candle_core::DType::U32, &self.device)?,
Tensor::new(vec![0u32], &self.device)?,
Tensor::zeros(text_length, candle_core::DType::U32, &self.device)?,
],
D::Minus1,
)?
.unsqueeze(0)?;
Some(mask)
};
let audio_feat = audio_feat.unsqueeze(0)?;
(Some(audio_feat), audio_mask)
} else {
(None, None)
};
let text_token = text_token.unsqueeze(0)?;
Ok((text_token, audio_feat, audio_mask))
}
pub fn processor_use_cache(
&self,
target_text: String,
prompt_cache: &HashMap<String, Tensor>,
tokenizer: &SingleChineseTokenizer,
) -> Result<(Tensor, Option<Tensor>, Option<Tensor>)> {
let (target_text_token, _) = tokenizer.encode_tensor(target_text, &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], &self.device)?;
let mut 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().unsqueeze(0)?)),
None => (0, None),
};
let audio_mask = if audio_length > 0 {
let text_pad_token =
Tensor::zeros(audio_length, candle_core::DType::U32, &self.device)?;
text_token = Tensor::cat(&[text_token, text_pad_token], D::Minus1)?;
let mask = Tensor::cat(
&[
Tensor::zeros(text_length, candle_core::DType::U32, &self.device)?,
Tensor::ones(audio_length, candle_core::DType::U32, &self.device)?,
],
D::Minus1,
)?
.unsqueeze(0)?;
Some(mask)
} else {
None
};
let text_token = text_token.unsqueeze(0)?;
Ok((text_token, audio_feat, audio_mask))
}
}
+9 -1
View File
@@ -366,7 +366,15 @@ pub fn get_audio_bytes_vec(path_str: &str) -> Result<Vec<u8>> {
let data = BASE64_STANDARD.decode(data)?; let data = BASE64_STANDARD.decode(data)?;
Ok(data) Ok(data)
} else { } else {
Err(anyhow::anyhow!("get audio path error {}", path_str)) let wave_u8 = path_str.as_bytes();
match get_audio_format_from_bytes(wave_u8) {
Ok(_) => Ok(wave_u8.to_vec()),
Err(e) => Err(anyhow::anyhow!(
"get audio path error {}, et_audio_format error: {}",
path_str,
e
)),
}
} }
} }
+82 -24
View File
@@ -1,5 +1,6 @@
use std::time::Instant; use std::time::Instant;
use aha::models::voxcpm_refact::generate::VoxCPMGenerateRefact;
use aha::params::chat::ChatCompletionParameters; use aha::params::chat::ChatCompletionParameters;
use aha::{ use aha::{
models::{ models::{
@@ -59,7 +60,7 @@ fn voxcpm1_5_use_message_generate() -> Result<()> {
#[test] #[test]
fn voxcpm1_5_generate() -> Result<()> { fn voxcpm1_5_generate() -> Result<()> {
// RUST_BACKTRACE=1 cargo test -F cuda voxcpm1_5_generate -r -- --nocapture // RUST_BACKTRACE=1 cargo test -F cuda --test test_voxcpm1_5 voxcpm1_5_generate -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!("{}/OpenBMB/VoxCPM1.5/", save_dir); let model_path = format!("{}/OpenBMB/VoxCPM1.5/", save_dir);
@@ -71,35 +72,37 @@ fn voxcpm1_5_generate() -> Result<()> {
let i_start = Instant::now(); let i_start = Instant::now();
// let generate = voxcpm_generate.generate_simple("太阳当空照,花儿对我笑,小鸟说早早早".to_string())?; // let generate = voxcpm_generate.generate_simple("太阳当空照,花儿对我笑,小鸟说早早早".to_string())?;
let generate = voxcpm_generate.inference( // let generate = voxcpm_generate.inference(
"老大爷我来啦,红红火火恍恍惚惚".to_string(), // "老大爷我来啦,红红火火恍恍惚惚".to_string(),
Some("啥子小师叔,打狗还要看主人,你再要继续,我就是你的对手".to_string()), // Some("啥子小师叔,打狗还要看主人,你再要继续,我就是你的对手".to_string()),
Some("file://./assets/audio/voice_01.wav".to_string()), // Some("file://./assets/audio/voice_01.wav".to_string()),
// Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()), // // Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()),
// Some("file://./assets/audio/voice_05.wav".to_string()), // // Some("file://./assets/audio/voice_05.wav".to_string()),
// 2,
// 4096,
// 10,
// 2.0,
// // false,
// 6.0,
// )?;
// 创建prompt_cache
voxcpm_generate.build_prompt_cache(
"啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string(),
"file://./assets/audio/voice_01.wav".to_string(),
)?;
// 使用prompt_cache生成语音
let generate = voxcpm_generate.generate_use_prompt_cache(
"太阳当空照,花儿对我笑,小鸟说早早早".to_string(),
2, 2,
4096, 100,
10, 10,
2.0, 2.0,
// false, false,
6.0, 6.0,
)?; )?;
// 创建prompt_cache std::thread::sleep(std::time::Duration::from_secs(2));
// let _ = voxcpm_generate.build_prompt_cache(
// "啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string(),
// "file://./assets/audio/voice_01.wav".to_string(),
// )?;
// // 使用prompt_cache生成语音
// let generate = voxcpm_generate.generate_use_prompt_cache(
// "太阳当空照,花儿对我笑,小鸟说早早早".to_string(),
// 2,
// 100,
// 10,
// 2.0,
// false,
// 6.0,
// )?;
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);
@@ -118,3 +121,58 @@ fn voxcpm1_5_tokenizer() -> Result<()> {
println!("ids: {:?}", ids); println!("ids: {:?}", ids);
Ok(()) Ok(())
} }
#[test]
fn voxcpm_refact_generate() -> Result<()> {
// RUST_BACKTRACE=1 cargo test -F cuda --test test_voxcpm1_5 voxcpm_refact_generate -r -- --nocapture
let save_dir =
aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?;
let model_path = format!("{}/OpenBMB/VoxCPM1.5/", save_dir);
let i_start = Instant::now();
let mut voxcpm_generate = VoxCPMGenerateRefact::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.inference(
// "VoxCPM is an innovative end-to-end TTS model from ModelBest, designed to generate highly realistic speech.".to_string(),
// Some("啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string()),
// Some("file://./assets/audio/voice_01.wav".to_string()),
// // Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()),
// // Some("file://./assets/audio/voice_05.wav".to_string()),
// 2,
// 100,
// 10,
// 2.0,
// // false,
// 6.0,
// )?;
// 创建prompt_cache
voxcpm_generate.build_prompt_cache(
"啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string(),
"file://./assets/audio/voice_01.wav".to_string(),
)?;
// 使用prompt_cache生成语音
let i_start = Instant::now();
let generate = voxcpm_generate.generate_use_prompt_cache(
"太阳当空照,花儿对我笑,小鸟说早早早".to_string(),
2,
100,
10,
2.0,
false,
6.0,
)?;
std::thread::sleep(std::time::Duration::from_secs(2));
let i_duration = i_start.elapsed();
println!("Time elapsed in generate is: {:?}", i_duration);
save_wav(
&generate,
"voxcpm.wav",
voxcpm_generate.sample_rate() as u32,
)?;
Ok(())
}