updata index tts

This commit is contained in:
jhqxxx
2026-02-14 15:52:30 +08:00
parent 7c832e0ce8
commit 2c34fc2d79
40 changed files with 2602 additions and 540 deletions
+100 -15
View File
@@ -1,27 +1,112 @@
// use std::io::Cursor;
use std::time::Instant;
use aha::utils::tensor_utils::interpolate_nearest_1d;
use anyhow::{Result, anyhow};
use candle_core::Tensor;
use sentencepiece::SentencePieceProcessor;
use std::fs::File;
// use symphonia::core::io::MediaSourceStream;
use std::io::{Read, Seek};
use std::{io::Cursor, time::Instant};
use aha::utils::{load_tensor_from_pt, tensor_utils::interpolate_nearest_1d};
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
use anyhow::{Result, anyhow};
use byteorder::{LittleEndian, ReadBytesExt};
use candle_core::{Shape, Tensor};
use sentencepiece::SentencePieceProcessor;
use zip::ZipArchive;
#[test]
fn messy_test() -> Result<()> {
// RUST_BACKTRACE=1 cargo test -F cuda messy_test -r -- --nocapture
let device = &candle_core::Device::Cpu;
let save_dir =
let save_dir: String =
aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?;
let model_path = format!("{}/IndexTeam/IndexTTS-2", save_dir);
let bpe_path = model_path.to_string() + "/bpe.model";
let tokenizer = SentencePieceProcessor::open(bpe_path)
.map_err(|e| anyhow!(format!("load bpe.model file error:{}", e)))?;
let tokens = tokenizer
.encode("你好啊")
.map_err(|e| anyhow!(format!("tokenizer encode error:{}", e)))?;
println!("tokens: {:?}", tokens);
let model_path = format!("{}/IndexTeam/IndexTTS-2/", save_dir);
let emo_matrix_path = model_path.clone() + "/feat2.pt";
let t_emo = load_tensor_from_pt(
&emo_matrix_path,
"feat2/data/0",
Shape::from_dims(&[73, 1280]),
&device,
)?;
println!("t_emo: {}", t_emo);
let skp_matrix_path = model_path + "/feat1.pt";
let t_skp = load_tensor_from_pt(
&skp_matrix_path,
"feat1/data/0",
Shape::from_dims(&[73, 192]),
&device,
)?;
println!("t_skp: {}", t_skp);
// let file = File::open(emo_matrix_path)?;
// let mut archive = ZipArchive::new(file)?;
// // 列出所有文件(调试用)
// for i in 0..archive.len() {
// let file = archive.by_index(i)?;
// println!("File: {} ({} bytes)", file.name(), file.size());
// }
// // 读取原始字节数据
// let mut data_file = archive.by_name("feat2/data/0")?;
// let mut buffer = Vec::new();
// data_file.read_to_end(&mut buffer)?;
// // 将字节转换为 f32 (little endian)
// let mut cursor = Cursor::new(buffer);
// let num_elements = 73 * 1280; // 93,440
// let mut data = Vec::with_capacity(num_elements);
// for _ in 0..num_elements {
// let val = cursor.read_f32::<LittleEndian>()?;
// data.push(val);
// }
// let t = Tensor::from_vec(data, (73, 1280), device)?;
// println!("t: {}", t);
// let message = r#"
// {
// "model": "index-tts2",
// "messages": [
// {
// "role": "user",
// "content": [
// {
// "type": "audio",
// "audio_url":
// {
// "url": "file:///home/jhq/Videos/voice_01.wav"
// }
// },
// {
// "type": "text",
// "text": "你好啊"
// }
// ]
// }
// ],
// "metadata": {"emo_vector": "[0, 0, 0, 0, 0, 0, 0.45, 0]"}
// }
// "#;
// let mes: ChatCompletionParameters = serde_json::from_str(message)?;
// if let Some(map) = &mes.metadata
// && let Some(emo_vector_str) = map.get("emo_vector")
// {
// match serde_json::from_str::<Vec<f32>>(emo_vector_str) {
// Ok(emo_vector) => {
// println!("Parsed emo_vector: {:?}", emo_vector);
// // 现在 emo_vector 是 Vec<f32>: [0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.45, 0.0]
// }
// Err(e) => {
// eprintln!("Failed to parse emo_vector: {}", e);
// }
// }
// }
// let save_dir =
// aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?;
// let model_path = format!("{}/IndexTeam/IndexTTS-2", save_dir);
// let bpe_path = model_path.to_string() + "/bpe.model";
// let tokenizer = SentencePieceProcessor::open(bpe_path)
// .map_err(|e| anyhow!(format!("load bpe.model file error:{}", e)))?;
// let tokens = tokenizer
// .encode("你好啊")
// .map_err(|e| anyhow!(format!("tokenizer encode error:{}", e)))?;
// println!("tokens: {:?}", tokens);
// let t = Tensor::arange(0.0f32, 40.0, device)?.broadcast_as((1, 40, 40))?;
// println!("t: {}", t);
// let i_start = Instant::now();
+12 -5
View File
@@ -1,8 +1,11 @@
use std::time::Instant;
use anyhow::Result;
use aha::models::index_tts2::{generate::IndexTTS2Generate, utils::download_index_tts2_need_model};
use aha::{
models::index_tts2::{generate::IndexTTS2Generate, utils::download_index_tts2_need_model},
utils::audio_utils::extract_and_save_audio_from_response,
};
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
use anyhow::Result;
#[tokio::test]
async fn index_tts2_generate() -> Result<()> {
@@ -24,10 +27,10 @@ async fn index_tts2_generate() -> Result<()> {
{
"url": "file:///home/jhq/Videos/voice_01.wav"
}
},
},
{
"type": "text",
"text": "你好啊"
"text": "你好啊,吃饭了吗"
}
]
}
@@ -42,7 +45,11 @@ async fn index_tts2_generate() -> Result<()> {
let i_start = Instant::now();
let generate = voxcpm_generate.generate(mes)?;
let save_path = extract_and_save_audio_from_response(&generate, "./")?;
for path in save_path {
println!("save audio: {}", path);
}
let i_duration = i_start.elapsed();
println!("Time elapsed in generate is: {:?}", i_duration);
Ok(())
}
}
+1 -1
View File
@@ -9,7 +9,7 @@ fn qwen3_asr_generate() -> Result<()> {
// RUST_BACKTRACE=1 cargo test -F cuda qwen3_asr_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!("{}/Qwen/Qwen3-ASR-0.6B/", save_dir); //Qwen/Qwen3-ASR-1.7B
let model_path = format!("{}/Qwen/Qwen3-ASR-0.6B/", save_dir); //Qwen/Qwen3-ASR-1.7B
let message = r#"
{
"model": "qwen3-asr",
+21 -10
View File
@@ -2,7 +2,11 @@ use std::collections::HashMap;
use aha::utils::{find_type_files, get_device, read_pth_tensor_info_cycle};
use anyhow::Result;
use candle_core::{Device, pickle::read_all_with_key, safetensors};
use candle_core::{
Device,
pickle::{read_all_with_key, read_pth_tensor_info},
safetensors,
};
use candle_nn::VarBuilder;
#[test]
@@ -203,19 +207,26 @@ fn index_tts2_weight() -> Result<()> {
// RUST_BACKTRACE=1 cargo test -F cuda index_tts2_weight -r -- --nocapture
let save_dir: String =
aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?;
let model_path = format!("{}/IndexTeam/IndexTTS-2/", save_dir);
let s2mel_path = model_path+ "/s2mel.pth";
let model_path = format!("{}/IndexTeam/IndexTTS-2/", save_dir);
let bigvgan_path = format!(
"{}/nv-community/bigvgan_v2_22khz_80band_256x/bigvgan_generator.pt",
save_dir
);
// let gpt_path = model_path+ "/gpt.pth";
// let spk_matrix_path = model_path+ "/feat1.pt";
// let s2mel_path = model_path+ "/s2mel.pth";
// let wac2vec2_path = model_path+ "/wav2vec2bert_stats.pt";
// let model_path = format!("{}/iic/speech_campplus_sv_zh-cn_16k-common/", save_dir);
// let campplus_path = model_path+ "/campplus_cn_common.bin";
// let model_list = find_type_files(&model_path, "safetensors")?;
let model_list = vec![s2mel_path];
// let mut dict_to_hashmap = HashMap::new();
// let mut dtype = candle_core::DType::F32;
let model_list = vec![bigvgan_path];
// // 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"))?;
// let dict = read_all_with_key(m, Some("net"))?;
let dict = read_pth_tensor_info_cycle(m, Some("net.cfm"))?;
let dict = read_all_with_key(m, Some("generator"))?;
// let dict = read_pth_tensor_info_cycle(m, Some("net.cfm"))?;
// dtype = dict[0].1.dtype();
for (k, v) in dict {
// if k.contains("model") {
@@ -230,9 +241,9 @@ fn index_tts2_weight() -> Result<()> {
// let model_list = vec![semantic_codec_path];
// for m in model_list {
// let weights = safetensors::load(m, &device)?;
// for (key, tensor) in weights.iter() {
// for (key, tensor) in weights.iter() {
// println!("=== {} === {:?}", key, tensor.shape());
// }
// }
Ok(())
}
}