delete some use
This commit is contained in:
@@ -453,7 +453,7 @@ impl CausalDecoder {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub fn forward(&self, x: &Tensor) -> Result<Tensor> {
|
pub fn forward(&self, x: &Tensor) -> Result<Tensor> {
|
||||||
let x = self.model0.forward(x)?;
|
let x = self.model0.forward(x)?;
|
||||||
let mut x = self.model1.forward(&x)?;
|
let mut x = self.model1.forward(&x)?;
|
||||||
for model_i in &self.model2_5 {
|
for model_i in &self.model2_5 {
|
||||||
x = model_i.forward(&x)?;
|
x = model_i.forward(&x)?;
|
||||||
|
|||||||
@@ -203,9 +203,9 @@ impl UnifiedCFM {
|
|||||||
estimator: VoxCPMLocDiT,
|
estimator: VoxCPMLocDiT,
|
||||||
mean_mode: bool,
|
mean_mode: bool,
|
||||||
) -> Result<Self> {
|
) -> Result<Self> {
|
||||||
let solver = cfm_params.solver;
|
// let solver = cfm_params.solver;
|
||||||
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,
|
||||||
@@ -305,9 +305,9 @@ impl UnifiedCFM {
|
|||||||
st_star = st_star.reshape(vec_shape)?;
|
st_star = st_star.reshape(vec_shape)?;
|
||||||
}
|
}
|
||||||
let cfg = cfg_dphi_dt.broadcast_mul(&st_star)?;
|
let cfg = cfg_dphi_dt.broadcast_mul(&st_star)?;
|
||||||
dphi_dt = cfg.add(&dphi_dt.sub(&cfg)?.affine(cfg_value, 0.0)?)?;
|
dphi_dt = cfg.add(&dphi_dt.sub(&cfg)?.affine(cfg_value, 0.0)?)?; // step步的预测噪声
|
||||||
}
|
}
|
||||||
x = x.broadcast_sub(&dphi_dt.broadcast_mul(&dt)?)?;
|
x = x.broadcast_sub(&dphi_dt.broadcast_mul(&dt)?)?; // 逐步去噪
|
||||||
t = t.sub(&dt)?;
|
t = t.sub(&dt)?;
|
||||||
sol.push(x.clone());
|
sol.push(x.clone());
|
||||||
if step < t_span_len - 1 {
|
if step < t_span_len - 1 {
|
||||||
@@ -598,10 +598,12 @@ impl VoxCPMModel {
|
|||||||
inference_timesteps,
|
inference_timesteps,
|
||||||
cfg_value,
|
cfg_value,
|
||||||
)?;
|
)?;
|
||||||
|
println!("laten_pred: {}", latent_pred);
|
||||||
let decode_audio = self
|
let decode_audio = self
|
||||||
.audio_vae
|
.audio_vae
|
||||||
.decode(&latent_pred.to_dtype(DType::F32)?)?
|
.decode(&latent_pred.to_dtype(DType::F32)?)?
|
||||||
.squeeze(1)?;
|
.squeeze(1)?;
|
||||||
|
println!("decode_audio: {}", decode_audio);
|
||||||
let decode_audio_len = decode_audio.dim(D::Minus1)? - 640 - 640;
|
let decode_audio_len = decode_audio.dim(D::Minus1)? - 640 - 640;
|
||||||
let decode_audio = decode_audio.narrow(D::Minus1, 640, decode_audio_len)?;
|
let decode_audio = decode_audio.narrow(D::Minus1, 640, decode_audio_len)?;
|
||||||
Ok(decode_audio)
|
Ok(decode_audio)
|
||||||
@@ -661,7 +663,6 @@ impl VoxCPMModel {
|
|||||||
let dit_hidden_2 = self.res_to_dit_proj.forward(&residual_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 = dit_hidden_1.add(&dit_hidden_2)?;
|
||||||
let cond = prefix_feat_cond.transpose(1, 2)?.contiguous()?;
|
let cond = prefix_feat_cond.transpose(1, 2)?.contiguous()?;
|
||||||
|
|
||||||
let pred_feat = self
|
let pred_feat = self
|
||||||
.feat_decoder
|
.feat_decoder
|
||||||
.forward(
|
.forward(
|
||||||
|
|||||||
@@ -1,5 +1,4 @@
|
|||||||
use anyhow::{Ok, Result, anyhow};
|
use anyhow::{Ok, Result, anyhow};
|
||||||
use candle_core::Tensor;
|
|
||||||
use tokenizers::Tokenizer;
|
use tokenizers::Tokenizer;
|
||||||
|
|
||||||
pub struct SingleChineseTokenizer {
|
pub struct SingleChineseTokenizer {
|
||||||
@@ -48,7 +47,7 @@ impl SingleChineseTokenizer {
|
|||||||
// println!("tokens: {:?}", tokens);
|
// println!("tokens: {:?}", tokens);
|
||||||
let mut split_character = Vec::new();
|
let mut split_character = Vec::new();
|
||||||
for token in tokens {
|
for token in tokens {
|
||||||
let clean_token = token.replace("▁", "to");
|
let clean_token = token.replace("▁", "");
|
||||||
if self.multichar_tokens.contains(&clean_token) {
|
if self.multichar_tokens.contains(&clean_token) {
|
||||||
let chars: Vec<String> = clean_token.chars().map(|c| c.to_string()).collect();
|
let chars: Vec<String> = clean_token.chars().map(|c| c.to_string()).collect();
|
||||||
split_character.extend(chars);
|
split_character.extend(chars);
|
||||||
|
|||||||
@@ -1,9 +1,8 @@
|
|||||||
use anyhow::{Result, anyhow};
|
use anyhow::{Result, anyhow};
|
||||||
use candle_core::{D, DType, Device, Tensor};
|
use candle_core::{D, Device, Tensor};
|
||||||
use candle_nn::{Conv1d, Conv1dConfig, Module, conv1d_no_bias};
|
use candle_nn::{Conv1d, Conv1dConfig, Module};
|
||||||
use hound::{SampleFormat, WavReader};
|
use hound::{SampleFormat, WavReader};
|
||||||
use rocket::futures::future::ok;
|
use num::integer::gcd;
|
||||||
|
|
||||||
use std::f64::consts::PI;
|
use std::f64::consts::PI;
|
||||||
use std::path::Path;
|
use std::path::Path;
|
||||||
|
|
||||||
@@ -14,10 +13,6 @@ pub enum ResamplingMethod {
|
|||||||
SincInterpKaiser,
|
SincInterpKaiser,
|
||||||
}
|
}
|
||||||
|
|
||||||
// 计算最大公约数
|
|
||||||
fn gcd(a: i64, b: i64) -> i64 {
|
|
||||||
if b == 0 { a } else { gcd(b, a % b) }
|
|
||||||
}
|
|
||||||
|
|
||||||
// 零阶修正贝塞尔函数 I0
|
// 零阶修正贝塞尔函数 I0
|
||||||
fn i0(x: f32) -> f32 {
|
fn i0(x: f32) -> f32 {
|
||||||
@@ -65,12 +60,12 @@ pub fn get_sinc_resample_kernel(
|
|||||||
|
|
||||||
let width_f = (lowpass_filter_width as f64) * (orig_freq as f64) / base_freq;
|
let width_f = (lowpass_filter_width as f64) * (orig_freq as f64) / base_freq;
|
||||||
let width = width_f.ceil() as i64;
|
let width = width_f.ceil() as i64;
|
||||||
// 创建索引数组 [1, 1, 2*width + orig_freq_reduced]
|
// 创建索引数组 [1, 1, 2*width + orig_freq]
|
||||||
let idx = Tensor::arange(-width as f32, (width + orig_freq) as f32, device)?
|
let idx = Tensor::arange(-width as f32, (width + orig_freq) as f32, device)?
|
||||||
.affine(1.0 / orig_freq as f64, 0.0)?
|
.affine(1.0 / orig_freq as f64, 0.0)?
|
||||||
.unsqueeze(0)?
|
.unsqueeze(0)?
|
||||||
.unsqueeze(0)?;
|
.unsqueeze(0)?;
|
||||||
// 创建时间数组 t [new_freq_reduced, 1, idx_len]
|
// 创建时间数组 t [new_freq, 1, idx_len]
|
||||||
let t = Tensor::arange_step(0.0, -new_freq as f32, -1.0, device)?
|
let t = Tensor::arange_step(0.0, -new_freq as f32, -1.0, device)?
|
||||||
.affine(1.0 / new_freq as f64, 0.0)?
|
.affine(1.0 / new_freq as f64, 0.0)?
|
||||||
.unsqueeze(D::Minus1)?
|
.unsqueeze(D::Minus1)?
|
||||||
@@ -270,7 +265,6 @@ pub fn load_audio<P: AsRef<Path>>(path: P, device: Device) -> Result<(Tensor, us
|
|||||||
&device,
|
&device,
|
||||||
)?
|
)?
|
||||||
.t()?;
|
.t()?;
|
||||||
// println!("audio channels: {}", spec.channels);
|
|
||||||
if spec.channels > 1 {
|
if spec.channels > 1 {
|
||||||
// 对channel通道求平均, channel维度变为1
|
// 对channel通道求平均, channel维度变为1
|
||||||
audio_tensor = audio_tensor.mean_keepdim(0)?;
|
audio_tensor = audio_tensor.mean_keepdim(0)?;
|
||||||
|
|||||||
@@ -1,6 +1,5 @@
|
|||||||
use aha::utils::audio_utils::{load_audio_with_resample};
|
use aha::utils::audio_utils::{load_audio_with_resample};
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
use candle_core::Tensor;
|
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn messy_test() -> Result<()> {
|
fn messy_test() -> Result<()> {
|
||||||
|
|||||||
+21
-21
@@ -22,28 +22,12 @@ fn voxcpm_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.generate(
|
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(),
|
|
||||||
)?;
|
|
||||||
// 使用prompt_cache生成语音
|
|
||||||
let generate = voxcpm_generate.generate_use_prompt_cache(
|
|
||||||
"太阳当空照,花儿对我笑,小鸟说早早早".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,
|
||||||
@@ -52,6 +36,22 @@ fn voxcpm_generate() -> Result<()> {
|
|||||||
6.0,
|
6.0,
|
||||||
)?;
|
)?;
|
||||||
|
|
||||||
|
// 创建prompt_cache
|
||||||
|
// let _ = voxcpm_generate.build_prompt_cache(
|
||||||
|
// "啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string(),
|
||||||
|
// "./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);
|
||||||
let _ = save_wav(&generate, "voxcpm.wav")?;
|
let _ = save_wav(&generate, "voxcpm.wav")?;
|
||||||
|
|||||||
Reference in New Issue
Block a user