use std::io::Write;
use audio_cpp::{Backend, ModelFamily, Registry, 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() < 3 {
eprintln!("用法: tts_offline <moss-tts-nano-100m-q8_0.gguf> <out.wav> [text]");
std::process::exit(1);
}
let model_path = &args[1];
let out_path = &args[2];
let text = args
.get(3)
.cloned()
.unwrap_or_else(|| "Hello from Rust and audio.cpp!".to_string());
let registry = Registry::new()?;
let families = registry.families()?;
println!("模型族: {families:?}");
if !families.iter().any(|f| f == ModelFamily::MossTtsNano.as_str()) {
eprintln!(
"警告: moss_tts_nano 未编译进引擎。请用 `--features custom-models`,\
并设置 AUDIOCPP_MODELS=moss_tts_nano 重新构建。"
);
}
let model = registry.load(model_path, Some(ModelFamily::MossTtsNano), None)?;
println!("模型加载成功: {model_path}");
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 request = format!(r#"{{"text":"{}"}}"#, text.replace('"', "\\\""));
println!("请求文本: {text}");
let result = session.run_offline(&request)?;
let audio = result
.audio_output
.as_ref()
.expect("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(())
}