diff --git a/data/2mix_30data/mixture_map.csv b/data/2mix_30data/mixture_map.csv deleted file mode 100644 index 76f63c3..0000000 --- a/data/2mix_30data/mixture_map.csv +++ /dev/null @@ -1,31 +0,0 @@ -mix_id,mix_file,spk1_file,spk2_file,spk1_name,spk2_name,overlap,snr_diff -m01,m01.wav,speaker10\speaker10_utt17.wav,speaker1\speaker1_utt16.wav,speaker10,speaker1,0.0,0 -m02,m02.wav,speaker9\speaker9_utt10.wav,speaker6\speaker6_utt1.wav,speaker9,speaker6,0.0,0 -m03,m03.wav,speaker8\speaker8_utt7.wav,speaker3\speaker3_utt3.wav,speaker8,speaker3,0.0,-5 -m04,m04.wav,speaker2\speaker2_utt19.wav,speaker6\speaker6_utt17.wav,speaker2,speaker6,0.0,0 -m05,m05.wav,speaker10\speaker10_utt12.wav,speaker6\speaker6_utt2.wav,speaker10,speaker6,0.3,-10 -m06,m06.wav,speaker7\speaker7_utt12.wav,speaker8\speaker8_utt20.wav,speaker7,speaker8,0.0,-10 -m07,m07.wav,speaker9\speaker9_utt11.wav,speaker3\speaker3_utt10.wav,speaker9,speaker3,0.6,0 -m08,m08.wav,speaker3\speaker3_utt20.wav,speaker10\speaker10_utt17.wav,speaker3,speaker10,0.3,-10 -m09,m09.wav,speaker5\speaker5_utt15.wav,speaker9\speaker9_utt17.wav,speaker5,speaker9,0.6,10 -m10,m10.wav,speaker8\speaker8_utt14.wav,speaker3\speaker3_utt4.wav,speaker8,speaker3,0.3,5 -m11,m11.wav,speaker1\speaker1_utt10.wav,speaker3\speaker3_utt19.wav,speaker1,speaker3,0.3,5 -m12,m12.wav,speaker9\speaker9_utt15.wav,speaker5\speaker5_utt5.wav,speaker9,speaker5,0.3,-10 -m13,m13.wav,speaker4\speaker4_utt16.wav,speaker2\speaker2_utt7.wav,speaker4,speaker2,0.6,-5 -m14,m14.wav,speaker5\speaker5_utt13.wav,speaker3\speaker3_utt6.wav,speaker5,speaker3,0.3,0 -m15,m15.wav,speaker2\speaker2_utt3.wav,speaker9\speaker9_utt9.wav,speaker2,speaker9,0.0,-5 -m16,m16.wav,speaker8\speaker8_utt7.wav,speaker4\speaker4_utt1.wav,speaker8,speaker4,0.6,10 -m17,m17.wav,speaker5\speaker5_utt18.wav,speaker10\speaker10_utt3.wav,speaker5,speaker10,0.0,5 -m18,m18.wav,speaker8\speaker8_utt6.wav,speaker2\speaker2_utt12.wav,speaker8,speaker2,0.6,5 -m19,m19.wav,speaker5\speaker5_utt7.wav,speaker2\speaker2_utt6.wav,speaker5,speaker2,0.0,-10 -m20,m20.wav,speaker1\speaker1_utt2.wav,speaker10\speaker10_utt18.wav,speaker1,speaker10,0.0,0 -m21,m21.wav,speaker10\speaker10_utt11.wav,speaker7\speaker7_utt7.wav,speaker10,speaker7,0.0,0 -m22,m22.wav,speaker4\speaker4_utt9.wav,speaker8\speaker8_utt3.wav,speaker4,speaker8,0.0,10 -m23,m23.wav,speaker6\speaker6_utt4.wav,speaker5\speaker5_utt6.wav,speaker6,speaker5,0.3,0 -m24,m24.wav,speaker10\speaker10_utt1.wav,speaker5\speaker5_utt8.wav,speaker10,speaker5,0.6,0 -m25,m25.wav,speaker10\speaker10_utt16.wav,speaker1\speaker1_utt11.wav,speaker10,speaker1,0.0,5 -m26,m26.wav,speaker4\speaker4_utt15.wav,speaker7\speaker7_utt7.wav,speaker4,speaker7,0.0,-10 -m27,m27.wav,speaker7\speaker7_utt15.wav,speaker6\speaker6_utt12.wav,speaker7,speaker6,0.0,-10 -m28,m28.wav,speaker6\speaker6_utt4.wav,speaker9\speaker9_utt10.wav,speaker6,speaker9,0.6,10 -m29,m29.wav,speaker6\speaker6_utt12.wav,speaker5\speaker5_utt16.wav,speaker6,speaker5,0.0,0 -m30,m30.wav,speaker6\speaker6_utt17.wav,speaker2\speaker2_utt4.wav,speaker6,speaker2,0.0,0 diff --git a/data/mix/mixture_map.csv b/data/mix/mixture_map.csv new file mode 100644 index 0000000..99e32cb --- /dev/null +++ b/data/mix/mixture_map.csv @@ -0,0 +1,51 @@ +mix_id,mix_file,src1,src2 +m01,m01.wav,speaker10\speaker10_17.wav,speaker1\speaker1_16.wav +m02,m02.wav,speaker10\speaker10_11.wav,speaker8\speaker8_08.wav +m03,m03.wav,speaker10\speaker10_16.wav,speaker3\speaker3_06.wav +m04,m04.wav,speaker3\speaker3_03.wav,speaker8\speaker8_16.wav +m05,m05.wav,speaker1\speaker1_03.wav,speaker2\speaker2_19.wav +m06,m06.wav,speaker5\speaker5_11.wav,speaker10\speaker10_20.wav +m07,m07.wav,speaker5\speaker5_10.wav,speaker4\speaker4_04.wav +m08,m08.wav,speaker6\speaker6_07.wav,speaker10\speaker10_18.wav +m09,m09.wav,speaker5\speaker5_11.wav,speaker3\speaker3_10.wav +m10,m10.wav,speaker4\speaker4_16.wav,speaker10\speaker10_12.wav +m11,m11.wav,speaker5\speaker5_02.wav,speaker2\speaker2_02.wav +m12,m12.wav,speaker10\speaker10_07.wav,speaker2\speaker2_16.wav +m13,m13.wav,speaker4\speaker4_16.wav,speaker8\speaker8_19.wav +m14,m14.wav,speaker1\speaker1_10.wav,speaker3\speaker3_19.wav +m15,m15.wav,speaker3\speaker3_15.wav,speaker5\speaker5_05.wav +m16,m16.wav,speaker7\speaker7_17.wav,speaker2\speaker2_13.wav +m17,m17.wav,speaker8\speaker8_08.wav,speaker4\speaker4_03.wav +m18,m18.wav,speaker5\speaker5_13.wav,speaker3\speaker3_06.wav +m19,m19.wav,speaker1\speaker1_13.wav,speaker10\speaker10_14.wav +m20,m20.wav,speaker9\speaker9_20.wav,speaker10\speaker10_20.wav +m21,m21.wav,speaker8\speaker8_07.wav,speaker4\speaker4_01.wav +m22,m22.wav,speaker8\speaker8_19.wav,speaker4\speaker4_12.wav +m23,m23.wav,speaker7\speaker7_17.wav,speaker1\speaker1_06.wav +m24,m24.wav,speaker10\speaker10_06.wav,speaker4\speaker4_09.wav +m25,m25.wav,speaker2\speaker2_06.wav,speaker8\speaker8_01.wav +m26,m26.wav,speaker1\speaker1_02.wav,speaker10\speaker10_18.wav +m27,m27.wav,speaker9\speaker9_11.wav,speaker10\speaker10_05.wav +m28,m28.wav,speaker8\speaker8_13.wav,speaker2\speaker2_05.wav +m29,m29.wav,speaker4\speaker4_09.wav,speaker8\speaker8_03.wav +m30,m30.wav,speaker8\speaker8_18.wav,speaker3\speaker3_20.wav +m31,m31.wav,speaker5\speaker5_06.wav,speaker7\speaker7_04.wav +m32,m32.wav,speaker10\speaker10_01.wav,speaker5\speaker5_08.wav +m33,m33.wav,speaker3\speaker3_11.wav,speaker1\speaker1_10.wav +m34,m34.wav,speaker1\speaker1_11.wav,speaker5\speaker5_06.wav +m35,m35.wav,speaker7\speaker7_07.wav,speaker3\speaker3_13.wav +m36,m36.wav,speaker9\speaker9_16.wav,speaker7\speaker7_05.wav +m37,m37.wav,speaker10\speaker10_03.wav,speaker9\speaker9_02.wav +m38,m38.wav,speaker1\speaker1_10.wav,speaker10\speaker10_20.wav +m39,m39.wav,speaker10\speaker10_15.wav,speaker3\speaker3_15.wav +m40,m40.wav,speaker6\speaker6_17.wav,speaker2\speaker2_04.wav +m41,m41.wav,speaker10\speaker10_07.wav,speaker7\speaker7_12.wav +m42,m42.wav,speaker1\speaker1_16.wav,speaker10\speaker10_14.wav +m43,m43.wav,speaker3\speaker3_10.wav,speaker6\speaker6_14.wav +m44,m44.wav,speaker4\speaker4_18.wav,speaker7\speaker7_03.wav +m45,m45.wav,speaker8\speaker8_13.wav,speaker7\speaker7_15.wav +m46,m46.wav,speaker1\speaker1_10.wav,speaker8\speaker8_19.wav +m47,m47.wav,speaker7\speaker7_06.wav,speaker8\speaker8_14.wav +m48,m48.wav,speaker10\speaker10_11.wav,speaker2\speaker2_09.wav +m49,m49.wav,speaker3\speaker3_12.wav,speaker6\speaker6_08.wav +m50,m50.wav,speaker1\speaker1_03.wav,speaker10\speaker10_08.wav diff --git a/examples/manual_add_speaker.py b/examples/manual_add_speaker.py new file mode 100644 index 0000000..5ab9c89 --- /dev/null +++ b/examples/manual_add_speaker.py @@ -0,0 +1,62 @@ +"""手動語者辨識測試腳本 + +此腳本允許一次針對某個 speaker 的多個檔案進行辨識, +只要給定資料夾路徑和索引列表即可自動跑完。 + +使用方式: + python -m examples.manual_add_speaker data/clean/speaker1 --indices 5 7 8 + # 會自動處理 speaker1_05.wav, speaker1_07.wav, speaker1_08.wav +""" + +import argparse +from pathlib import Path +from modules.identification import SpeakerIdentifier + + +def main() -> None: + """解析參數並依序執行語者辨識。""" + parser = argparse.ArgumentParser( + description="針對某個 speaker 資料夾內指定編號的檔案,執行語者辨識" + ) + parser.add_argument( + "speaker_dir", + help="speaker 資料夾路徑(如 data/clean/speaker1)" + ) + parser.add_argument( + "--indices", "-i", + required=True, + nargs="+", + type=int, + metavar="N", + help="要處理的檔案編號列表(不含前綴零),如 5 7 8" + ) + args = parser.parse_args() + + speaker_path = Path(args.speaker_dir) + if not speaker_path.is_dir(): + parser.error(f"{speaker_path!r} 不是一個有效的目錄") + + speaker_name = speaker_path.name # e.g. "speaker1" + identifier = SpeakerIdentifier() + + for idx in args.indices: + # 兩位數補零 + filename = f"{speaker_name}_{idx:02d}.wav" + audio_path = speaker_path / filename + + print(f"\n🔍 處理:{audio_path}") + if not audio_path.is_file(): + print(f"⚠️ 檔案不存在,跳過:{audio_path}") + continue + + result = identifier.process_audio_file(str(audio_path)) + if result: + speaker_id, speaker_label, distance = result + print(f"▶️ 語者: {speaker_label} (UUID: {speaker_id}) 相似度 {distance:.3f}") + else: + print("⚠️ 處理失敗或無法辨識") + +if __name__ == "__main__": + main() + +#python -m examples.manual_add_speaker data/clean/speaker1 -i 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 \ No newline at end of file diff --git a/examples/pipeline_eval.py b/examples/pipeline_eval.py new file mode 100644 index 0000000..883c9fc --- /dev/null +++ b/examples/pipeline_eval.py @@ -0,0 +1,395 @@ +import argparse +import csv +from dataclasses import dataclass +from pathlib import Path +from typing import Dict, List, Tuple +import time +import json +import gc + +# import numpy as np +import torch +import torchaudio +import tempfile + +from pipelines.orchestrator import init_pipeline_modules,run_pipeline_file +from modules.separation.separator import AudioSeparator +from modules.asr.text_utils import compute_cer, normalize_zh +# from modules.identification.VID_identify_v5 import SpeakerIdentifier +# from modules.asr.whisper_asr import WhisperASR + + +@dataclass +class SourceInfo: + path: Path + transcript: str | None + mix_sdr: float | None = None + sep_sdr: float | None = None + delta_sdr: float | None = None + clean_wave: torch.Tensor | None = None + + +def load_truth_map(path: Path) -> Dict[str, str]: + mapping: Dict[str, str] = {} + with path.open("r", encoding="utf-8") as f: + next(f) # skip header + for line in f: + line = line.strip() + if not line: + continue + fname, txt = line.split(",", 1) + mapping[fname] = txt + return mapping + + +def load_mixture_map(path: Path) -> List[Tuple[str, str, str, str]]: + rows: List[Tuple[str, str, str, str]] = [] + with path.open("r", encoding="utf-8") as f: + reader = csv.DictReader(f) + for row in reader: + rows.append((row["mix_id"], row["mix_file"], row["src1"], row["src2"])) + return rows + +#計算SI-SDR 可同時計算分離前和分離後的SDR +def si_sdr(est: torch.Tensor, ref: torch.Tensor) -> float: + """Scale-Invariant SDR for single-channel signals.""" + min_len = min(len(est), len(ref)) + est = est[:min_len] + ref = ref[:min_len] + + est = est - est.mean() + ref = ref - ref.mean() + s_target = torch.dot(est, ref) * ref / torch.dot(ref, ref) + e_noise = est - s_target + return 10 * torch.log10((s_target.pow(2).sum()) / (e_noise.pow(2).sum())).item() +#取檔案並統計SI-SDR +def compute_baseline_sisdr( + mix_path: Path, + src1_path: Path, + src2_path: Path, + separator: AudioSeparator, + sep_paths: List[Path] | None = None, +) -> Dict[str, float]: + """ + 計算 baseline 與分離後的 SI-SDR,以及 ΔSI-SDR。 + + mix_path : 混音檔 (.wav) + src1_path : 來源 1 (乾淨) + src2_path : 來源 2 (乾淨) + separator : 你初始化好的 AudioSeparator 實例 + sep_paths : 【可選】已存在的分離後 wav 路徑清單 + 若 None,函式會臨時再跑一次 separator 取得分離音 + return : dict,含 4 個欄位 + """ + # --- 讀取 waveforms -------------------------------------------------- + def _load(wav: Path) -> torch.Tensor: + wav_tensor, _ = torchaudio.load(wav) + return wav_tensor.squeeze(0) # [1, T] -> [T] + + mix_wave = _load(mix_path) + src1_wave = _load(src1_path) + src2_wave = _load(src2_path) + + # --- baseline:混音 vs 兩個乾淨源 ----------------------------------- + sdr_mix_src1 = si_sdr(mix_wave, src1_wave) + sdr_mix_src2 = si_sdr(mix_wave, src2_wave) + + # --- 取得分離後音檔 --------------------------------------------------- + if sep_paths is None: + # 如果呼叫端沒給,就臨時再分一次(省事但較耗時) + tmp_dir = Path(tempfile.mkdtemp(prefix="eval_sep_")) + # 注意:separator 需要 tensor 在正確 device + mix_wave_device = mix_wave.to(separator.device).unsqueeze(0) # [T] → [1, T] + sep_paths = [ + Path(p) for p, *_ in + separator.separate_and_save(mix_wave_device, tmp_dir.as_posix(), segment_index=0) + ] + + # --- 分離後 vs 乾淨源:取「最佳對應」即可 ----------------------------- + def _best_sdr(ref_wave: torch.Tensor) -> float: + best = -1e9 + for p in sep_paths: + est_wave = _load(p) + best = max(best, si_sdr(est_wave, ref_wave)) + return best + + sdr_sep_src1 = _best_sdr(src1_wave) + sdr_sep_src2 = _best_sdr(src2_wave) + + # --- ΔSI-SDR --------------------------------------------------------- + delta1 = sdr_sep_src1 - sdr_mix_src1 + delta2 = sdr_sep_src2 - sdr_mix_src2 + + return { + "si_sdr_src1": sdr_sep_src1, + "si_sdr_src2": sdr_sep_src2, + "delta_si_sdr_src1": delta1, + "delta_si_sdr_src2": delta2, + } +#spkID +def compute_accuracy(pred_speakers: List[str], true_speakers: List[str]) -> float: + """ + pred_speakers: pipeline 分離後所有段落的預測 speaker id (e.g. ["spk3","spk10"]) + true_speakers: mixture_map 定義的兩位原始 speaker id (e.g. ["spk10","spk1"]) + 回傳值: 正確數 / 2 -> 0.0, 0.5, or 1.0 + """ + matched = len(set(pred_speakers) & set(true_speakers)) + return matched / len(true_speakers) + +NUM_MAP = {'0':'零','1':'一','2':'二','3':'三','4':'四', + '5':'五','6':'六','7':'七','8':'八','9':'九'} + +def normalize_numbers_to_zh(text: str) -> str: + return "".join(NUM_MAP.get(ch, ch) for ch in text) + +def load_config_and_data() -> Tuple[argparse.Namespace, Path, Path, Dict[str, str], List[Tuple[str, str, str, str]]]: + parser = argparse.ArgumentParser(description="Run speech pipeline and evaluate") + parser.add_argument("--mix-dir", default="data/mix", help="mix路徑") + parser.add_argument("--clean-dir", default="data/clean", help="clean路徑") + parser.add_argument("--truth-map", default="data/truth_map.csv") + parser.add_argument("--test-list", default="data/mix/mixture_map.csv", help="batch 測試清單 CSV (mix_id,mix_file,src1,src2)") + parser.add_argument("--out", default="work_output/pipeline_results.csv") + args = parser.parse_args() + + mix_dir = Path(args.mix_dir) + clean_dir = Path(args.clean_dir) + truth_map = load_truth_map(Path(args.truth_map)) + mixture_rows = load_mixture_map(Path(args.test_list)) + Path(args.out).parent.mkdir(parents=True, exist_ok=True) + + return args, mix_dir, clean_dir, truth_map, mixture_rows + +def prepare_results_structure() -> tuple[list[str], list[dict]]: + columns = [ + "mix_id", "mix_file", + "sep_time", "sid_time", "asr_time", "total_time", + "si_sdr_src1", "si_sdr_src2", + "delta_si_sdr_src1", "delta_si_sdr_src2", + "accuracy", + "cer1", "cer2", "avg_conf", + "ref_text1", "pred_text1", "ref_text2", "pred_text2", + ] + return columns, [] + +def process_one_mixture( + mix_path: str, + mix_id: str = None, + sep=None, spk=None, asr=None, + clean_dir: Path = None, + mixture_rows: List[Tuple[str, str, str, str]] = None, + truth_map: Dict[str, str] = None, +) -> dict: + """ + 處理一組混音音檔,跑完整的 pipeline,並回傳所有可取得的指標。 + + mix_path: 混音音檔完整路徑 + mix_id: 選填的檔案名稱(不含副檔名),預設用檔名 + return: 一個 dict,包含 pipeline 的所有可觀察數據 + """ + # 使用目前時間計算總耗時 + overall_start = time.perf_counter() + mix_wav = Path(mix_path) + + if mix_id is None: + mix_id = Path(mix_path).stem # 取檔名當作 id + + # 🌀 呼叫主流程跑完整 pipeline + try: + bundle, _ , stats = run_pipeline_file(mix_path,3, sep=sep, spk=spk, asr=asr) + except Exception as e: + print(f"[ERROR] 處理檔案 {mix_path} 時出錯:{e}") + return { + "mix_id": mix_id, + "mix_file": mix_path, + "error": str(e) + } + + # 📊 統計與整理各階段資訊 + recog_texts = [s.get("text", "") for s in bundle if s.get("text")] + confidences = [s.get("confidence", 0.0) for s in bundle if s.get("text")] + + full_text = " ".join(recog_texts) + avg_conf = sum(confidences) / len(confidences) if confidences else 0.0 + + result = { + "mix_id": mix_id, + "mix_file": mix_path, + "seg_count": len(bundle), + "sep_time": round(stats.get("separate", 0.0), 3), + "sid_time": round(stats.get("speaker", 0.0), 3), + "asr_time": round(stats.get("asr", 0.0), 3), + "total_time": round(stats.get("total", 0.0), 3), + "pipeline_time": round(stats.get("pipeline_total", 0.0), 3), + "avg_conf": round(avg_conf, 4), + "predicted_text": full_text, + "segments": bundle, + "overall_time": round(time.perf_counter() - overall_start, 3), + } + + # ③ 加上 baseline SI-SDR 計算 + # 解析 src1 / src2 清音路徑 + if clean_dir and mixture_rows: + row = next((r for r in mixture_rows if r[0] == mix_id), None) + if row: + src1 = clean_dir / row[2] + src2 = clean_dir / row[3] + sep_paths = [Path(s["path"]) for s in bundle if "path" in s] + sisdr_metrics = compute_baseline_sisdr(mix_wav, src1, src2, sep, sep_paths) + result.update(sisdr_metrics) + + # --- 取出 true speaker IDs from mixture_map row --- + # row[2] = "speaker10/…", row[3] = "speaker1/…" + # row[2] = "speaker10/speaker10_17.wav" + # Path(row[2]).parent.name → "speaker10" + # .replace("speaker", "spk") → "spk10" + + true_spk1 = Path(row[2]).parent.name.replace("speaker", "spk") + true_spk2 = Path(row[3]).parent.name.replace("speaker", "spk") + true_speakers = [true_spk1, true_spk2] + # print(f"🔍 真實語者:{true_speakers}") + + # --- 取出 pipeline 預測的所有 speaker IDs --- + pred_speakers = [seg.get("speaker") for seg in bundle] + # print(f"🔍 預測語者:{pred_speakers}") + # --- 計算 accuracy --- + acc = compute_accuracy(pred_speakers, true_speakers) + # print(f"🔍 語者辨識準確率:{acc:.2f}") + result["accuracy"] = acc + + # === CER & ref/pred text for each speaker === # + # 1) 地取每個 source 的正確文字 + ref_fname1 = Path(src1).name # e.g. "speaker10_17.wav" + ref_fname2 = Path(src2).name + raw_ref1 = truth_map.get(ref_fname1, "") + raw_ref2 = truth_map.get(ref_fname2, "") + + # 2) Normalize (去標點、空格、統一繁體) + ref_norm1 = normalize_zh(raw_ref1) + ref_norm2 = normalize_zh(raw_ref2) + # print(f"🔍 正規化文字1:{ref_norm1}" + # f" 正規化文字2:{ref_norm2}") + + # 3) 直接抓 bundle 前兩段文字,不理會 speaker id + pred_texts = [seg.get("text", "") for seg in bundle if seg.get("text")] + if len(pred_texts) < 2: # 若只抓到 1 段,就補空 + pred_texts.append("") + predA, predB = pred_texts[:2] # A = 第一段, B = 第二段 + + normA = normalize_numbers_to_zh(normalize_zh(predA)) + normB = normalize_numbers_to_zh(normalize_zh(predB)) + # print(f"🔍 預測文字1:{pred_norm1}" + # f" 預測文字2:{pred_norm2}") + # 4) 交叉計算 4 個 CER + cA1 = compute_cer(ref_norm1, normA) if ref_norm1 else None + cB2 = compute_cer(ref_norm2, normB) if ref_norm2 else None + cA2 = compute_cer(ref_norm2, normA) if ref_norm2 else None + cB1 = compute_cer(ref_norm1, normB) if ref_norm1 else None + # print(f"🔍 CER1:{cer1:.4f} CER2:{cer2:.4f}") + + # 5) 選「加總最小」的配對 + if (cA1 or 0) + (cB2 or 0) <= (cA2 or 0) + (cB1 or 0): + final_pred1, final_pred2 = normA, normB + cer1, cer2 = cA1, cB2 + else: + final_pred1, final_pred2 = normB, normA + cer1, cer2 = cB1, cA2 + # 5) 寫入結果 + result.update({ + "ref_text1": ref_norm1, + "pred_text1": final_pred1, + "cer1": cer1, + "ref_text2": ref_norm2, + "pred_text2": final_pred2, + "cer2": cer2, + }) + torch.cuda.empty_cache() + gc.collect() + else: + print(f"[WARN] mix_id {mix_id} not found in mixture_map") + + + + return result + +def run_and_evaluate_pipeline( + mixture_csv: Path, + mix_dir: Path, + clean_dir: Path, + truth_map: Dict[str, str], + sep, spk, asr, + output_csv: Path, +): + test_rows = load_mixture_map(mixture_csv) # ① 讀取清單 + cols, _ = prepare_results_structure() + + with output_csv.open("w", newline="", encoding="utf-8") as fout: + writer = csv.DictWriter(fout, fieldnames=cols) + writer.writeheader() + + for mix_id, mix_file, *_ in test_rows: # ② 逐筆跑 + mix_path = mix_dir / mix_file + print(f"👉 處理 {mix_id} → {mix_path}") + + res = process_one_mixture( + str(mix_path), mix_id, + sep=sep, spk=spk, asr=asr, + clean_dir=clean_dir, + mixture_rows=test_rows, + truth_map=truth_map + ) + + # ③ 如果 pipeline 出錯就跳過,不寫進 CSV + if "error" in res: + print(f"🚨 跳過 {mix_id},原因:{res['error']}") + continue + + # ④ 只留下指定欄位 + row = {k: res.get(k, "") for k in cols} + writer.writerow(row) + + print(f"✅ 全部完成!結果已寫入 {output_csv}") + +def main(): + # 1️⃣ 清空 GPU 快取 + torch.cuda.empty_cache() + gc.collect() + + # 2️⃣ 載入所有參數與資料 + # --mix-dir、--clean-dir、--truth-map、--test-list、--out 都已設預設 + args, mix_dir, clean_dir, truth_map, mixture_rows = load_config_and_data() + + # 3️⃣ 初始化模型 + sep, spk, asr, _ = init_pipeline_modules() + + run_and_evaluate_pipeline( + mixture_csv=Path(args.test_list), + mix_dir=mix_dir, + clean_dir=clean_dir, + truth_map=truth_map, + sep=sep, spk=spk, asr=asr, + output_csv=Path(args.out), + ) + + # # ④ 測試一組混音 + # test_mix_id = "m05" # 你可以自訂一個 ID + # test_mix_path = "data/mix/m05.wav" # ← 改成你實際有的音檔路徑! + + # print(f"👉 處理測試音檔:{test_mix_path}") + # result = process_one_mixture( + # test_mix_path, + # test_mix_id, + # sep=sep, + # spk=spk, + # asr=asr, + # clean_dir=clean_dir, + # mixture_rows=mixture_rows, + # truth_map=truth_map + # ) + # # 寫入 JSON 檔 + # with open(f"{test_mix_id}_result.json", "w", encoding="utf-8") as f: + # json.dump(result, f, indent=2, ensure_ascii=False) + # print(f"💾 已儲存 JSON 檔至:{test_mix_id}_result.json") + + +if __name__ == "__main__": + main() diff --git a/examples/predefined_speaker_builder.py b/examples/predefined_speaker_builder.py new file mode 100644 index 0000000..7c0ac7f --- /dev/null +++ b/examples/predefined_speaker_builder.py @@ -0,0 +1,113 @@ +"""預先建立語者並使用指定音檔平均建立聲紋。 + +此腳本示範如何在程式碼中預先指定音檔索引, +依序將其平均成一個聲紋,並建立對應語者。 +""" + +from __future__ import annotations + +import os +from typing import Dict, List + +import numpy as np + +from modules.database import DatabaseService # ← 新增 +from modules.identification import SpeakerIdentifier + +# === 語者設定 ============================================================== +# key: 新語者名稱 +# value: {"folder": 資料夾路徑, "indices": [音檔索引列表]} +SPEAKER_CONFIG: Dict[str, Dict[str, List[int]]] = { + "spk1": { + "folder": "data/clean/speaker1", + "indices": [1, 4, 5, 7, 8, 9, 12, 14, 15, 17, 18, 19, 20], + }, + "spk2": { + "folder": "data/clean/speaker2", + "indices": [1, 3, 7, 8, 10, 11, 12, 14, 15, 17, 18, 20], + }, + "spk3": { + "folder": "data/clean/speaker3", + "indices": [1, 2, 4, 5, 7, 8, 9, 14, 16, 17, 18], + }, + "spk4": { + "folder": "data/clean/speaker4", + "indices": [2, 5, 6, 7, 8, 10, 11, 13, 14, 15, 17, 19, 20], + }, + "spk5": { + "folder": "data/clean/speaker5", + "indices": [1, 3, 4, 7, 9, 12, 14, 15, 16, 17, 18, 19, 20], + }, + "spk6": { + "folder": "data/clean/speaker6", + "indices": [1, 2, 3, 4, 5, 6, 9, 10, 11, 12, 13, 15, 16, 18, 19, 20], + }, + "spk7": { + "folder": "data/clean/speaker7", + "indices": [1, 2, 8, 9, 10, 11, 13, 14, 16, 18, 19, 20], + }, + "spk8": { + "folder": "data/clean/speaker8", + "indices": [2, 4, 5, 6, 9, 10, 11, 12, 15, 17, 20], + }, + "spk9": { + "folder": "data/clean/speaker9", + "indices": [1, 3, 4, 5, 6, 7, 8, 9, 10, 12, 13, 14, 15, 17, 18, 19], + }, + "spk10": { + "folder": "data/clean/speaker10", + "indices": [2, 4, 9, 10, 13, 19], + }, + # 範例:若要建立更多語者,可在此處加入設定 + # "n4": { + # "folder": "data/clean/speaker4", + # "indices": [1, 3, 5, 7], + # }, +} + +# 初始化識別器與資料庫 +identifier = SpeakerIdentifier() +db = DatabaseService() # ← 用新的物件 + + +def build_speaker(name: str, folder: str, indices: List[int]) -> None: + """依索引平均音檔嵌入並建立語者。 + + Args: + name: 新語者名稱。 + folder: 音檔所在資料夾。 + indices: 需要使用的音檔索引(不含前置零)。 + """ + base = os.path.basename(folder) + embeddings = [] + + for idx in indices: + file_path = os.path.join(folder, f"{base}_{idx:02d}.wav") + if not os.path.exists(file_path): + print(f"音檔 {file_path} 不存在,已跳過。") + continue + emb = identifier.audio_processor.extract_embedding(file_path) + embeddings.append(emb) + + if not embeddings: + print(f"語者 {name} 沒有有效音檔,跳過建立。") + return + + avg_embedding = np.mean(np.stack(embeddings), axis=0) + + # 建立語者 + first_file = os.path.join(folder, f"{base}_{indices[0]:02d}.wav") + speaker_uuid = db.create_speaker(full_name=name,first_audio=first_file) + + # 建立平均聲紋 + db.create_voiceprint( + speaker_uuid, + avg_embedding, + audio_source="avg_of_indices", + ) + print(f"已建立語者 {name} (UUID: {speaker_uuid}),使用 {len(embeddings)} 個音檔。") + + +if __name__ == "__main__": + for sp_name, cfg in SPEAKER_CONFIG.items(): + build_speaker(sp_name, cfg["folder"], cfg["indices"]) \ No newline at end of file diff --git a/examples/stats_summarizer.py b/examples/stats_summarizer.py new file mode 100644 index 0000000..33d94b2 --- /dev/null +++ b/examples/stats_summarizer.py @@ -0,0 +1,97 @@ +# stats_summarizer.py +import pandas as pd +import numpy as np + +AUDIO_LEN_SEC = 6.0 # 固定音檔長度(秒) +LOW_CONF_THRESHOLD = 0.6 # Low-confidence 門檻 + +def compute_summary(df: pd.DataFrame) -> dict: + out = {} + + # --- CER:把 cer1 + cer2 併在一起統計 --- + cer_all = np.concatenate([df["cer1"].values, df["cer2"].values]) + out["cer_mean"] = float(np.mean(cer_all)) + out["cer_median"] = float(np.median(cer_all)) + out["cer_p90"] = float(np.percentile(cer_all, 90)) + out["cer_std"] = float(np.std(cer_all, ddof=0)) + + # --- SI-SDR:每列只取「正的那個」;兩個都正就取最大;兩個都<=0 就跳過 --- + sisdr_pick = [] + for s1, s2 in zip(df["si_sdr_src1"], df["si_sdr_src2"]): + candidates = [x for x in (s1, s2) if x > 0] + if len(candidates) == 1: + sisdr_pick.append(candidates[0]) + elif len(candidates) > 1: + sisdr_pick.append(max(candidates)) + # 若兩個都 <= 0,則忽略(不列入平均) + out["sisdr_mean_pos_only"] = float(np.mean(sisdr_pick)) if sisdr_pick else float("nan") + out["sisdr_count_used"] = int(len(sisdr_pick)) + out["sisdr_rows_total"] = int(len(df)) + + # --- ΔSI-SDR:直接把兩欄併在一起平均(你的規則) --- + delta_all = np.concatenate([df["delta_si_sdr_src1"].values, df["delta_si_sdr_src2"].values]) + out["delta_sisdr_mean"] = float(np.mean(delta_all)) + out["delta_sisdr_median"] = float(np.median(delta_all)) + out["delta_sisdr_p90"] = float(np.percentile(delta_all, 90)) + out["delta_sisdr_std"] = float(np.std(delta_all, ddof=0)) + + # --- Confidence --- + out["avg_conf_mean"] = float(df["avg_conf"].mean()) + out["low_conf_pct@0.6"] = float((df["avg_conf"] < LOW_CONF_THRESHOLD).mean() * 100.0) + + # --- Accuracy 分桶 + 平均 --- + acc = df["accuracy"] + out["accuracy_mean"] = float(acc.mean()) + out["acc_1_pct"] = float((acc == 1).mean() * 100.0) + out["acc_0_5_pct"] = float((acc == 0.5).mean() * 100.0) + out["acc_0_pct"] = float((acc == 0).mean() * 100.0) + + # --- RTF 與耗時占比 --- + rtf = df["total_time"] / AUDIO_LEN_SEC + out["rtf_mean"] = float(rtf.mean()) + out["rtf_median"] = float(rtf.median()) + out["rtf_p90"] = float(np.percentile(rtf, 90)) + out["rtf_std"] = float(rtf.std(ddof=0)) + + total = df["total_time"].replace({0: np.nan}) + out["share_sep_pct_mean"] = float(((df["sep_time"] / total) * 100.0).mean()) + out["share_sid_pct_mean"] = float(((df["sid_time"] / total) * 100.0).mean()) + out["share_asr_pct_mean"] = float(((df["asr_time"] / total) * 100.0).mean()) + + return out + +def summarize_csv(in_csv: str, out_csv: str) -> None: + df = pd.read_csv(in_csv) + s = compute_summary(df) + summary_df = pd.DataFrame({ + "Metric": [ + "CER mean", "CER median", "CER p90", "CER std", + "SISDR (pos-only) mean", "SISDR used rows / total", + "ΔSI-SDR mean", "ΔSI-SDR median", "ΔSI-SDR p90", "ΔSI-SDR std", + "Avg confidence", "Low-confidence % @0.6", + "Accuracy mean", "Acc=1 %", "Acc=0.5 %", "Acc=0 %", + "RTF mean", "RTF median", "RTF p90", "RTF std", + "Share: Separation % (mean)", + "Share: SpeakerID % (mean)", + "Share: ASR % (mean)", + ], + "Value": [ + s["cer_mean"], s["cer_median"], s["cer_p90"], s["cer_std"], + s["sisdr_mean_pos_only"], f'{s["sisdr_count_used"]} / {s["sisdr_rows_total"]}', + s["delta_sisdr_mean"], s["delta_sisdr_median"], s["delta_sisdr_p90"], s["delta_sisdr_std"], + s["avg_conf_mean"], s["low_conf_pct@0.6"], + s["accuracy_mean"], s["acc_1_pct"], s["acc_0_5_pct"], s["acc_0_pct"], + s["rtf_mean"], s["rtf_median"], s["rtf_p90"], s["rtf_std"], + s["share_sep_pct_mean"], s["share_sid_pct_mean"], s["share_asr_pct_mean"], + ] + }) + summary_df.to_csv(out_csv, index=False) + +if __name__ == "__main__": + import argparse + p = argparse.ArgumentParser() + p.add_argument("--in_csv", required=True, help="Path to pipeline_results.csv") + p.add_argument("--out_csv", default="pipeline_summary.csv", help="Where to write the summary CSV") + args = p.parse_args() + summarize_csv(args.in_csv, args.out_csv) + print(f"Saved summary to: {args.out_csv}") diff --git a/modules/asr/whisper_asr.py b/modules/asr/whisper_asr.py index 55572fd..ab16440 100644 --- a/modules/asr/whisper_asr.py +++ b/modules/asr/whisper_asr.py @@ -20,7 +20,7 @@ class WhisperASR: text, confidence, words = asr.transcribe("path/to/audio.wav") """ - def __init__(self, model_name: str = None, gpu: bool = False, beam: int = None, lang: str = "auto"): + def __init__(self, model_name: str = None, gpu: bool = False, beam: int = None, lang: str = "zh"): self.gpu = gpu self.beam = beam if beam is not None else DEFAULT_WHISPER_BEAM_SIZE self.lang = lang diff --git a/modules/separation/separator.py b/modules/separation/separator.py index e08e2b9..f9a1ac0 100644 --- a/modules/separation/separator.py +++ b/modules/separation/separator.py @@ -99,11 +99,6 @@ """ import os - -# 修復 SVML 錯誤:在導入 PyTorch 之前設定環境變數 -os.environ["MKL_DISABLE_FAST_MM"] = "1" -os.environ["KMP_DUPLICATE_LIB_OK"] = "TRUE" - import numpy as np import torch import torchaudio @@ -120,8 +115,6 @@ from scipy import signal from scipy.ndimage import uniform_filter1d from enum import Enum -from sklearn.cluster import DBSCAN # type: ignore -import librosa # type: ignore # 修改模型載入方式 try: @@ -180,7 +173,7 @@ class SeparationModel(Enum): TARGET_RATE = AUDIO_TARGET_RATE WINDOW_SIZE = AUDIO_WINDOW_SIZE OVERLAP = AUDIO_OVERLAP -DEVICE_INDEX = 2 +DEVICE_INDEX = None # 處理參數(從配置讀取) MIN_ENERGY_THRESHOLD = AUDIO_MIN_ENERGY_THRESHOLD @@ -294,12 +287,6 @@ def __init__(self, model_type: SeparationModel = DEFAULT_MODEL, enable_noise_red self.enable_noise_reduction = enable_noise_reduction self.snr_threshold = snr_threshold - # 語者偵測相關參數 - 進一步降低為更靈敏的設定 - self.vad_threshold = 0.1 # 進一步降低語音活動檢測閾值(原0.08) - self.speaker_energy_threshold = 0.1 # 進一步降低說話者能量閾值(原0.15) - self.silence_threshold = 0.002 # 進一步降低靜音檢測閾值(原0.002) - self.min_speech_duration = 0.8 # 進一步降低最小語音持續時間(原0.3秒) - logger.info(f"使用設備: {self.device}") logger.info(f"模型類型: {model_type.value}") logger.info(f"載入模型: {self.model_config['model_name']}") @@ -913,246 +900,9 @@ def _log_final_statistics(self): f"跳過: {stats['segments_skipped']}, " f"錯誤: {stats['errors']}") - def detect_speaker_count(self, audio_tensor: torch.Tensor) -> int: - """ - 偵測音訊中的說話者數量 - - Args: - audio_tensor: 輸入音訊張量 [channels, samples] 或 [batch, channels, samples] - - Returns: - int: 偵測到的說話者數量 - """ - try: - # 確保音訊格式正確 - if len(audio_tensor.shape) == 3: - audio_data = audio_tensor[0, 0, :].cpu().numpy() - elif len(audio_tensor.shape) == 2: - audio_data = audio_tensor[0, :].cpu().numpy() - else: - audio_data = audio_tensor.cpu().numpy() - - # 0. 基本靜音檢測 - 進一步放寬條件 - audio_rms = np.sqrt(np.mean(audio_data ** 2)) - audio_max = np.max(np.abs(audio_data)) - audio_std = np.std(audio_data) - - logger.info(f"音訊統計 - RMS: {audio_rms:.6f}, Max: {audio_max:.6f}, STD: {audio_std:.6f}, 靜音閾值: {self.silence_threshold:.6f}") - - # 更寬鬆的靜音檢測:只要有一個指標超過閾值就認為可能有語音 - is_silent = (audio_rms < self.silence_threshold and - audio_max < self.silence_threshold * 4 and - audio_std < self.silence_threshold * 1.5) - - if is_silent: - logger.info(f"判定為靜音 - 所有指標都低於閾值") - return 0 - - logger.info(f"通過靜音檢測,進行語音活動偵測...") - - # 1. 簡化的語音活動檢測 (VAD) - frame_length = int(TARGET_RATE * 0.025) # 25ms 幀 - hop_length = int(TARGET_RATE * 0.010) # 10ms 跳躍 - - # 計算短時能量 - frames = librosa.util.frame(audio_data, frame_length=frame_length, hop_length=hop_length) - energy = np.sum(frames ** 2, axis=0) - - # 正規化能量 - if np.max(energy) > 0: - energy = energy / np.max(energy) - - # 簡化的VAD:主要基於能量檢測 - energy_active = energy > self.vad_threshold - - # 額外的過零率檢測作為輔助(但不是必需的) - try: - zero_crossings = [] - for i in range(min(frames.shape[1], 100)): # 限制處理數量以提高速度 - frame = frames[:, i] - zcr = np.sum(np.diff(np.sign(frame)) != 0) / len(frame) - zero_crossings.append(zcr) - zero_crossings = np.array(zero_crossings) - - # 如果有合理的過零率,結合使用;否則只用能量 - if len(zero_crossings) > 0: - zcr_active = (zero_crossings > 0.005) & (zero_crossings < 0.8) # 非常寬鬆的範圍 - if np.sum(zcr_active) > len(zcr_active) * 0.05: # 如果至少5%的幀有合理過零率 - combined_active = energy_active[:len(zcr_active)] & zcr_active - if np.sum(combined_active) > np.sum(energy_active) * 0.3: # 如果結合檢測結果不會過度降低 - vad_frames = np.zeros_like(energy_active, dtype=bool) - vad_frames[:len(combined_active)] = combined_active - else: - vad_frames = energy_active - else: - vad_frames = energy_active - else: - vad_frames = energy_active - except Exception as e: - logger.debug(f"過零率檢測失敗,使用純能量檢測: {e}") - vad_frames = energy_active - - logger.info(f"VAD結果 - 總幀數: {len(vad_frames)}, 活動幀數: {np.sum(vad_frames)}, 比例: {np.sum(vad_frames)/len(vad_frames):.3f}") - - # 2. 更寬鬆的語音區段檢查 - if np.any(vad_frames): - total_speech_ratio = np.sum(vad_frames) / len(vad_frames) - - # 大幅降低語音活動比例要求 - if total_speech_ratio < 0.02: # 只要2%的幀有活動就可能是語音 - logger.info(f"語音活動比例太低: {total_speech_ratio:.3f}") - return 0 - - # 檢查連續區段(但要求更寬鬆) - diff = np.diff(np.concatenate(([False], vad_frames, [False])).astype(int)) - starts = np.where(diff == 1)[0] - ends = np.where(diff == -1)[0] - - if len(starts) > 0 and len(ends) > 0: - speech_durations = (ends - starts) * hop_length / TARGET_RATE - max_duration = np.max(speech_durations) if len(speech_durations) > 0 else 0 - - logger.info(f"語音區段分析 - 區段數: {len(starts)}, 最長持續時間: {max_duration:.2f}秒, 最小要求: {self.min_speech_duration:.2f}秒") - - # 如果有任何區段超過最小要求,或者總活動比例足夠 - if max_duration >= self.min_speech_duration or total_speech_ratio > 0.05: - logger.info("通過語音區段檢查") - else: - logger.info(f"語音區段太短且活動比例不足") - return 0 - else: - # 如果無法計算區段,檢查總體活動 - if total_speech_ratio < 0.03: - logger.info(f"無法計算區段且活動比例不足: {total_speech_ratio:.3f}") - return 0 - else: - logger.info("未檢測到任何語音活動") - return 0 - - # 3. 如果通過了所有檢查,嘗試進行說話者數量估計 - try: - # 簡化的MFCC分析 - mfccs = librosa.feature.mfcc( - y=audio_data, - sr=TARGET_RATE, - n_mfcc=13, - hop_length=hop_length, - n_fft=frame_length*2 - ) - - # 確保維度匹配 - min_frames = min(mfccs.shape[1], len(vad_frames)) - mfccs = mfccs[:, :min_frames] - vad_frames = vad_frames[:min_frames] - - if np.any(vad_frames): - active_mfccs = mfccs[:, vad_frames] - - if active_mfccs.shape[1] < 5: # 進一步降低要求 - logger.info(f"有效語音幀數很少 ({active_mfccs.shape[1]}),直接判定為單一語者") - return 1 - - # 簡化的說話者數量估計 - features = active_mfccs.T - audio_duration = len(audio_data) / TARGET_RATE - - # 對於短音訊,直接判定為單一語者 - if audio_duration < 3.0 or features.shape[0] < 20: - logger.info(f"短音訊 ({audio_duration:.2f}秒) 或特徵少 ({features.shape[0]}),判定為單一語者") - return 1 - - # 嘗試聚類分析(但如果失敗就回傳1) - try: - clustering = DBSCAN( - eps=1.2, # 更大的聚類半徑 - min_samples=max(3, int(features.shape[0] * 0.1)) - ).fit(features) - - unique_labels = set(clustering.labels_) - if -1 in unique_labels: - unique_labels.remove(-1) - - detected_speakers = len(unique_labels) - - if detected_speakers == 0: - detected_speakers = 1 # 後備方案 - - logger.info(f"聚類分析結果: {detected_speakers} 位說話者") - - except Exception as e: - logger.info(f"聚類分析失敗,預設為單一語者: {e}") - detected_speakers = 1 - - else: - detected_speakers = 1 - - except Exception as e: - logger.info(f"MFCC分析失敗,預設為單一語者: {e}") - detected_speakers = 1 - - # 4. 最終限制和驗證 - detected_speakers = min(detected_speakers, self.num_speakers) - detected_speakers = max(detected_speakers, 1) # 既然通過了語音檢測,至少是1個說話者 - - logger.info(f"最終偵測結果: {detected_speakers} 位說話者") - return detected_speakers - - except Exception as e: - logger.warning(f"語者數量偵測失敗,使用後備方案: {e}") - return self._fallback_speaker_detection(audio_data if 'audio_data' in locals() else audio_tensor.cpu().numpy()) - - def _fallback_speaker_detection(self, audio_data: np.ndarray) -> int: - """ - 後備的語者數量偵測方法,基於簡化的能量分析 - - Args: - audio_data: 音訊數據 - - Returns: - int: 估計的說話者數量 - """ - try: - # 非常寬鬆的靜音檢測 - audio_rms = np.sqrt(np.mean(audio_data ** 2)) - audio_max = np.max(np.abs(audio_data)) - - logger.info(f"後備檢測 - RMS: {audio_rms:.6f}, Max: {audio_max:.6f}") - - # 只要有一個指標超過很低的閾值就認為有語音 - if audio_rms < self.silence_threshold * 0.3 and audio_max < self.silence_threshold: - logger.info("後備檢測:判定為靜音") - return 0 - - # 如果通過靜音檢測,至少回傳1個說話者 - logger.info("後備檢測:判定為有語音活動") - return 1 - - except Exception as e: - logger.warning(f"後備語者偵測也失敗: {e}") - # 終極後備方案:只要音訊不是完全靜音就認為有1個說話者 - try: - if np.any(np.abs(audio_data) > 0.0001): # 極低的閾值 - logger.info("終極後備:偵測到非零音訊") - return 1 - else: - logger.info("終極後備:音訊完全靜音") - return 0 - except: - logger.info("所有檢測都失敗,預設回傳1") - return 1 # 最保險的選擇 - def separate_and_save(self, audio_tensor, output_dir, segment_index): """分離並儲存音訊,並回傳 (path, start, end) 列表。""" try: - # 新增:在分離前先偵測說話者數量 - detected_speakers = self.detect_speaker_count(audio_tensor) - logger.info(f"片段 {segment_index} - 偵測到 {detected_speakers} 位說話者") - - # 如果沒有偵測到說話者,跳過處理 - if detected_speakers == 0: - logger.info(f"片段 {segment_index} - 未偵測到說話者,跳過處理") - return [] - # 初始化累計時間戳 current_t0 = getattr(self, "_current_t0", 0.0) results = [] # 用來收 (path, start, end) @@ -1210,11 +960,7 @@ def separate_and_save(self, audio_tensor, output_dir, segment_index): saved_count = 0 start_time = current_t0 - - # 根據偵測到的說話者數量限制輸出 - effective_speakers = min(detected_speakers, num_speakers, self.num_speakers) - - for i in range(effective_speakers): + for i in range(min(num_speakers, self.num_speakers)): try: if speaker_dim == 1: speaker_audio = enhanced_separated[0, i, :].cpu() @@ -1261,7 +1007,7 @@ def separate_and_save(self, audio_tensor, output_dir, segment_index): logger.warning(f"儲存語者 {i+1} 失敗: {e}") if saved_count > 0: - logger.info(f"片段 {segment_index} 完成,實際儲存 {saved_count}/{effective_speakers} 個檔案") + logger.info(f"片段 {segment_index} 完成,儲存 {saved_count} 個檔案") # 更新累計時間到下一段 current_t0 += seg_duration @@ -1282,15 +1028,6 @@ def separate_and_save(self, audio_tensor, output_dir, segment_index): def separate_and_identify(self, audio_tensor: torch.Tensor, output_dir: str, segment_index: int) -> None: """分離音訊並直接進行語音識別,可選擇是否儲存音訊檔案""" try: - # 新增:在分離前先偵測說話者數量 - detected_speakers = self.detect_speaker_count(audio_tensor) - logger.info(f"片段 {segment_index} - 偵測到 {detected_speakers} 位說話者") - - # 如果沒有偵測到說話者,跳過處理 - if detected_speakers == 0: - logger.info(f"片段 {segment_index} - 未偵測到說話者,跳過處理") - return [] - audio_files = [] audio_streams = [] diff --git a/pipelines/orchestrator.py b/pipelines/orchestrator.py index 6b66ab9..2776fec 100644 --- a/pipelines/orchestrator.py +++ b/pipelines/orchestrator.py @@ -33,30 +33,33 @@ logger.info(" Device: %s", torch.cuda.get_device_name(0)) # ---------- 1. GPU/CPU 設備選擇 ---------- -current_cuda_device = CUDA_DEVICE_INDEX # 建立本地變數避免修改全域變數 - -if FORCE_CPU: - use_gpu = False - logger.info("🔧 FORCE_CPU=true,強制使用 CPU") -else: - use_gpu = torch.cuda.is_available() - if use_gpu: - # 檢查指定的設備是否存在 - if current_cuda_device < torch.cuda.device_count(): - torch.cuda.set_device(current_cuda_device) - logger.info(f"🎯 設定 CUDA 設備索引: {current_cuda_device}") - logger.info(f" 使用設備: {torch.cuda.get_device_name(current_cuda_device)}") - else: - logger.warning(f"⚠️ CUDA 設備索引 {current_cuda_device} 不存在,使用預設設備 0") - current_cuda_device = 0 - torch.cuda.set_device(current_cuda_device) # 確實設定設備 0 - logger.info(f" 已設定為設備 0: {torch.cuda.get_device_name(0)}") - -logger.info(f"🚀 使用設備: {'cuda:' + str(current_cuda_device) if use_gpu else 'cpu'}") +def init_pipeline_modules(): + """初始化 sep / spk / asr,考慮 CUDA 設定,並回傳模組實例們""" + current_cuda_device = CUDA_DEVICE_INDEX # 建立本地變數避免修改全域變數 -sep = AudioSeparator() -spk = SpeakerIdentifier() -asr = WhisperASR(model_name=DEFAULT_WHISPER_MODEL, gpu=use_gpu, beam=DEFAULT_WHISPER_BEAM_SIZE) + if FORCE_CPU: + use_gpu = False + logger.info("🔧 FORCE_CPU=true,強制使用 CPU") + else: + use_gpu = torch.cuda.is_available() + if use_gpu: + if current_cuda_device < torch.cuda.device_count(): + torch.cuda.set_device(current_cuda_device) + logger.info(f"🎯 設定 CUDA 設備索引: {current_cuda_device}") + logger.info(f" 使用設備: {torch.cuda.get_device_name(current_cuda_device)}") + else: + logger.warning(f"⚠️ CUDA 設備索引 {current_cuda_device} 不存在,使用預設設備 0") + current_cuda_device = 0 + torch.cuda.set_device(current_cuda_device) + logger.info(f" 已設定為設備 0: {torch.cuda.get_device_name(0)}") + + logger.info(f"🚀 使用設備: {'cuda:' + str(current_cuda_device) if use_gpu else 'cpu'}") + + sep = AudioSeparator() + spk = SpeakerIdentifier() + asr = WhisperASR(model_name=DEFAULT_WHISPER_MODEL, gpu=use_gpu, beam=DEFAULT_WHISPER_BEAM_SIZE) + + return sep, spk, asr,use_gpu def _timed_call(func, *args): t0 = time.perf_counter() @@ -64,7 +67,7 @@ def _timed_call(func, *args): return res, time.perf_counter() - t0 -def process_segment(seg_path: str, t0: float, t1: float) -> dict: +def process_segment(seg_path: str, t0: float, t1: float,sep=None, spk=None, asr=None) -> dict: """Process a single separated segment.""" logger.info( f"🔧 執行緒 {threading.get_ident()} 處理 ({t0:.2f}-{t1:.2f}) → {os.path.basename(seg_path)}" @@ -110,6 +113,7 @@ def process_segment(seg_path: str, t0: float, t1: float) -> dict: "words": adjusted_words, "spk_time": spk_time, "asr_time": asr_time, + "path": seg_path } @@ -127,8 +131,19 @@ def make_pretty(seg: dict) -> dict: } -def run_pipeline_file(raw_wav: str, max_workers: int = 3): - """Run pipeline on an existing wav file.""" +def run_pipeline_file(raw_wav: str, max_workers: int = 3, sep=None, spk=None, asr=None): + """Run pipeline on an existing wav file. + + If the separator / speaker identifier / ASR modules are not provided, + they will be initialized automatically. This preserves backwards + compatibility for callers that import :func:`run_pipeline_file` directly + without using :func:`init_pipeline_modules` first (e.g. older API code). + """ + + # Allow legacy usage where modules are not injected explicitly. + if sep is None or spk is None or asr is None: + sep, spk, asr, _ = init_pipeline_modules() + total_start = time.perf_counter() waveform, sr = torchaudio.load(raw_wav) @@ -148,9 +163,22 @@ def run_pipeline_file(raw_wav: str, max_workers: int = 3): logger.info(f"⏱ 分離耗時 {sep_end - sep_start:.3f}s, 共 {len(segments)} 段") # 2) 多執行緒處理所有段 + torch.cuda.empty_cache() # ★ 釋放分離占用的 VRAM logger.info(f"🔄 處理 {len(segments)} 段... (max_workers={max_workers})") with ThreadPoolExecutor(max_workers=max_workers) as ex: - bundle = [r for r in ex.map(lambda s: process_segment(*s), segments) if r] + bundle = [ + r for r in ex.map( + lambda s: process_segment( + s[0], s[1], s[2], + sep=sep, # ← 仍傳 sep,但此時顯存已空閒 + spk=spk, + asr=asr # Whisper v3 現在才真正占用 GPU + ), + segments + ) + if r + ] + spk_time = max((s.get("spk_time", 0.0) for s in bundle), default=0.0) asr_time = max((s.get("asr_time", 0.0) for s in bundle), default=0.0) @@ -226,13 +254,18 @@ def load_truth_map(path: str) -> dict[str, str]: def run_pipeline_dir( dir_path: str, truth_map_path: str = "truth_map.txt", - max_workers: int = 3, + max_workers: int = 3, sep=None, spk=None, asr=None ) -> str: """ 批次處理資料夾內所有音檔,輸出: - summary.tsv:檔案級統計 + 段落詳情 - asr_report.tsv:ASR 指標 (avg_conf, WER, CER) """ + + # Lazily initialize modules for legacy callers that do not provide them. + if sep is None or spk is None or asr is None: + sep, spk, asr, _ = init_pipeline_modules() + timestamp = dt.now().strftime("%Y%m%d_%H%M%S") out_dir = pathlib.Path("work_output") / f"batch_{timestamp}" out_dir.mkdir(parents=True, exist_ok=True) @@ -252,7 +285,7 @@ def run_pipeline_dir( file_results: list[tuple[int, Path, dict, list[dict]]] = [] for idx, audio in enumerate(sorted(audio_files), start=1): logger.info(f"===== 處理檔案 {audio.name} ({idx}/{len(audio_files)}) =====") - segments, pretty, stats = run_pipeline_file(str(audio), max_workers) + segments, pretty, stats = run_pipeline_file(str(audio), max_workers , sep=sep, spk=spk, asr=asr) # 用 truth_map 覆寫 WER/CER gt = truth_map.get(audio.name) @@ -310,9 +343,14 @@ def run_pipeline_stream( queue_out: "queue.Queue[dict] | None" = None, stop_event: threading.Event | None = None, in_bytes_queue: "queue.Queue[bytes] | None" = None, + sep=None, spk=None, asr=None ): """串流模式:每 chunk_secs 做一次分離/識別/ASR。""" + # Initialize modules when not supplied to maintain backward compatibility + if sep is None or spk is None or asr is None: + sep, spk, asr, _ = init_pipeline_modules() + total_start = time.perf_counter() out_root = Path("stream_output") / dt.now().strftime("%Y%m%d_%H%M%S") out_root.mkdir(parents=True, exist_ok=True) @@ -340,7 +378,7 @@ def process_chunk(raw_bytes: bytes, idx: int): speaker_results: list[dict] = [] for sp_idx, wav_path in enumerate(speaker_paths, 1): - res = process_segment(str(wav_path), t0, t1) + res = process_segment(str(wav_path), t0, t1, sep=sep, spk=spk, asr=asr) if not res["text"].strip() or res["confidence"] < 0.1: continue res["speaker_index"] = sp_idx @@ -481,6 +519,7 @@ def recorder_from_mic(): # ───────────────────────── CLI ───────────────────────── def main(): + sep, spk, asr, use_gpu = init_pipeline_modules() parser = argparse.ArgumentParser(description="Speech pipeline") sub = parser.add_subparsers(dest="mode", required=True) @@ -513,15 +552,17 @@ def main(): args = parser.parse_args() # 用 CLI 覆蓋 ASR 設定 - global asr + # 如果命令行 override 模型,就重新拿一个新的 asr asr = WhisperASR(model_name=args.model, gpu=use_gpu, beam=args.beam) if args.mode == "file": - run_pipeline_file(args.path, args.workers) + run_pipeline_file(args.path, + args.workers, + sep=sep, spk=spk, asr=asr) elif args.mode == "stream": - run_pipeline_stream(chunk_secs=args.chunk, max_workers=args.workers) + run_pipeline_stream(chunk_secs=args.chunk, max_workers=args.workers, sep=sep, spk=spk, asr=asr) elif args.mode == "dir": - run_pipeline_dir(args.path, truth_map_path=args.truth_map, max_workers=args.workers) + run_pipeline_dir(args.path, truth_map_path=args.truth_map, max_workers=args.workers, sep=sep, spk=spk, asr=asr) if __name__ == "__main__": diff --git a/utils/constants.py b/utils/constants.py index 3e99434..c8fe129 100644 --- a/utils/constants.py +++ b/utils/constants.py @@ -4,9 +4,9 @@ """ # 語者識別閾值 (演算法核心參數,經過實驗調校) -THRESHOLD_LOW = 0.26 # 過於相似,不更新向量 -THRESHOLD_UPDATE = 0.34 # 相似度足夠,更新向量 -THRESHOLD_NEW = 0.385 # 超過此值視為新語者 +THRESHOLD_LOW = 0.2 # 過於相似,不更新向量 +THRESHOLD_UPDATE = 0.33 # 相似度足夠,更新向量 +THRESHOLD_NEW = 0.37 # 超過此值視為新語者 # 音訊處理固定參數 (技術規格要求) AUDIO_SAMPLE_RATE = 16000 # SpeechBrain 模型要求的取樣率