use std::fs::File;
use std::path::Path;
use audio_cpp::{AudioInput, Backend, ModelFamily, Registry, Request, RunMode, TaskKind, WavAudio};
fn decode_any_audio(path: &str) -> Result<WavAudio, String> {
use symphonia::core::codecs::audio::AudioDecoderOptions;
use symphonia::core::formats::probe::Hint;
use symphonia::core::formats::{FormatOptions, TrackType};
use symphonia::core::io::{MediaSourceStream, MediaSourceStreamOptions};
use symphonia::core::meta::MetadataOptions;
let file = File::open(path).map_err(|e| format!("打开文件失败: {e}"))?;
let mss = MediaSourceStream::new(Box::new(file), MediaSourceStreamOptions::default());
let mut format = symphonia::default::get_probe()
.probe(
&Hint::new(),
mss,
FormatOptions::default(),
MetadataOptions::default(),
)
.map_err(|e| format!("探测格式失败: {e}"))?;
let track = format
.default_track(TrackType::Audio)
.ok_or_else(|| "文件中没有音频轨".to_string())?;
let codec_params = track
.codec_params
.as_ref()
.and_then(|c| c.audio())
.ok_or_else(|| "缺少音频编码参数".to_string())?;
let sample_rate = codec_params.sample_rate.unwrap_or(0) as i32;
let channels = codec_params
.channels
.as_ref()
.map(|c| c.count() as i32)
.unwrap_or(0);
let track_id = track.id;
let mut decoder = symphonia::default::get_codecs()
.make_audio_decoder(codec_params, &AudioDecoderOptions::default())
.map_err(|e| format!("创建解码器失败: {e}"))?;
let mut samples: Vec<f32> = Vec::new();
while let Some(packet) = format
.next_packet()
.map_err(|e| format!("读取数据包失败: {e}"))?
{
if packet.track_id != track_id {
continue;
}
let decoded = decoder
.decode(&packet)
.map_err(|e| format!("解码失败: {e}"))?;
let n = decoded.samples_interleaved();
if n == 0 {
continue;
}
let start = samples.len();
samples.resize(start + n, 0.0);
decoded.copy_to_slice_interleaved(&mut samples[start..]);
}
Ok(WavAudio {
sample_rate,
channels,
samples,
})
}
fn main() -> Result<(), audio_cpp::Error> {
let args: Vec<String> = std::env::args().collect();
if args.len() < 3 {
eprintln!("用法: load_any_audio <权重.safetensors> <input 任意格式> [family_hint]");
std::process::exit(1);
}
let model_path = &args[1];
let audio_path = &args[2];
let family_hint = args.get(3).map(String::as_str).map(ModelFamily::from);
let wav = decode_any_audio(audio_path).map_err(audio_cpp::Error::Other)?;
println!(
"解码成功: {} 采样率={}Hz 声道={} 采样数={}",
Path::new(audio_path)
.file_name()
.and_then(|n| n.to_str())
.unwrap_or(audio_path),
wav.sample_rate,
wav.channels,
wav.samples.len()
);
let registry = Registry::new()?;
let model = registry.load(model_path, family_hint.clone(), None)?;
println!("模型加载成功: {model_path} family_hint={family_hint:?}");
let session = model.create_task_session(
TaskKind::Vad,
RunMode::Offline,
Backend::Cpu,
0, 4, None,
)?;
let threshold_key = if session.family() == "marblenet_vad" {
"threshold"
} else {
"vad_threshold"
};
let result =
session.run_offline(Request::vad(AudioInput::Buffer(wav)).option(threshold_key, 0.5))?;
println!("=== 语音片段 ===");
for seg in &result.speech_segments {
println!(
" {}..{} 置信度={}",
seg.span.start_sample, seg.span.end_sample, seg.confidence
);
}
Ok(())
}