#![allow(clippy::unwrap_used, clippy::expect_used)]
use std::io::Write;
use audio_cpp::{Backend, ModelFamily, Registry, Request, RunMode, TaskKind};
fn write_wav_pcm16(
path: &str,
samples: &[f32],
sample_rate: i32,
channels: u16,
) -> std::io::Result<()> {
let bytes_per_sample = 2u32; let block_align = bytes_per_sample * channels as u32;
let byte_rate = sample_rate as u32 * block_align;
let data_len = samples.len() as u32 * bytes_per_sample;
let mut f = std::fs::File::create(path)?;
f.write_all(b"RIFF")?;
f.write_all(&(36u32 + data_len).to_le_bytes())?;
f.write_all(b"WAVE")?;
f.write_all(b"fmt ")?;
f.write_all(&16u32.to_le_bytes())?;
f.write_all(&1u16.to_le_bytes())?; f.write_all(&channels.to_le_bytes())?;
f.write_all(&(sample_rate as u32).to_le_bytes())?;
f.write_all(&byte_rate.to_le_bytes())?;
f.write_all(&(block_align as u16).to_le_bytes())?;
f.write_all(&(bytes_per_sample as u16 * 8).to_le_bytes())?; f.write_all(b"data")?;
f.write_all(&data_len.to_le_bytes())?;
for &s in samples {
let v = (s.clamp(-1.0, 1.0) * 32767.0) as i16;
f.write_all(&v.to_le_bytes())?;
}
Ok(())
}
fn main() -> Result<(), audio_cpp::Error> {
let args: Vec<String> = std::env::args().collect();
if args.len() < 4 {
eprintln!(
"用法: tts_offline_qwen3 <model.gguf> <reference.wav> <reference_text> <out.wav> [text]"
);
std::process::exit(1);
}
let model_path = &args[1];
let reference_path = &args[2];
let reference_text = &args[3];
let out_path = &args[4];
let text = args
.get(5)
.cloned()
.unwrap_or_else(|| "Hello from Rust and Qwen3 TTS!".to_string());
let registry = Registry::new()?;
println!("模型族: {:?}", registry.families()?);
let model = registry.load(model_path, Some(ModelFamily::Qwen3Tts), None)?;
println!("模型加载成功: {model_path}");
println!("元数据: {:?}", model.metadata()?);
let session = model.create_task_session(
TaskKind::Tts,
RunMode::Offline,
Backend::Cpu,
0, 4, None,
)?;
println!(
"会话: family={} task={} mode={}",
session.family(),
session.task_kind(),
session.run_mode()
);
let result = session.run_offline(
Request::tts(&text)
.reference(reference_path)
.reference_text(reference_text),
)?;
println!("参考音频: {reference_path}");
println!("参考文本: {reference_text}");
println!("合成文本: {text}");
let audio = result
.audio_output
.as_ref()
.expect("Qwen3 TTS 应返回 audio_output");
let samples = audio
.samples
.as_deref()
.expect("audio_output 应携带 samples 数据");
if samples.is_empty() {
return Err(audio_cpp::Error::Ffi("合成音频为空".to_string()));
}
let channels = audio.channels.max(1) as u16;
write_wav_pcm16(out_path, samples, audio.sample_rate, channels)
.map_err(|e| audio_cpp::Error::Ffi(format!("写 WAV 失败: {e}")))?;
println!(
"已写入 {out_path}: {}Hz {}ch {} 采样({} 秒)",
audio.sample_rate,
audio.channels,
samples.len(),
samples.len() as f64 / audio.sample_rate as f64
);
Ok(())
}