698 lines
24 KiB
Python
698 lines
24 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""
|
||
说话人分离模块
|
||
基于 CAM++ 的说话人分离,用于多说话人音频分割
|
||
"""
|
||
|
||
from loguru import logger
|
||
import numpy as np
|
||
import librosa
|
||
import soundfile as sf
|
||
import tempfile
|
||
import os
|
||
import threading
|
||
from typing import Any, List, Mapping, Optional, Sequence, cast
|
||
from dataclasses import dataclass
|
||
|
||
import torch
|
||
|
||
from ..core.config import settings
|
||
from ..core.exceptions import DefaultServerErrorException
|
||
|
||
# 全局 CAM++ pipeline 缓存(单例)
|
||
_global_diarization_pipeline: Any | None = None
|
||
_diarization_pipeline_lock = threading.Lock()
|
||
_diarization_inference_semaphore = threading.BoundedSemaphore(1)
|
||
|
||
|
||
@dataclass
|
||
class SpeakerSegment:
|
||
"""说话人分段信息"""
|
||
|
||
start_ms: int
|
||
end_ms: int
|
||
speaker_id: str
|
||
audio_data: Optional[np.ndarray] = None
|
||
temp_file: Optional[str] = None
|
||
|
||
@property
|
||
def start_sec(self) -> float:
|
||
return self.start_ms / 1000.0
|
||
|
||
@property
|
||
def end_sec(self) -> float:
|
||
return self.end_ms / 1000.0
|
||
|
||
@property
|
||
def duration_ms(self) -> int:
|
||
return self.end_ms - self.start_ms
|
||
|
||
@property
|
||
def duration_sec(self) -> float:
|
||
return self.duration_ms / 1000.0
|
||
|
||
|
||
def _resolve_modelscope_device() -> str:
|
||
"""根据配置和硬件自动选择 modelscope pipeline 设备
|
||
"""
|
||
from ..core.device import detect_device
|
||
|
||
return detect_device(settings.DEVICE)
|
||
|
||
|
||
def _move_pipeline_model_to_device(pipeline_instance: Any, modelscope_device: str) -> None:
|
||
"""将 pipeline 的底层模型迁移到目标设备。"""
|
||
if hasattr(pipeline_instance, "device_name"):
|
||
pipeline_instance.device_name = modelscope_device
|
||
model = getattr(pipeline_instance, "model", None)
|
||
if model is not None and hasattr(model, "to"):
|
||
pipeline_instance.model = model.to(modelscope_device)
|
||
|
||
|
||
def _create_modelscope_pipeline(
|
||
*,
|
||
task: Any,
|
||
model: str,
|
||
modelscope_device: str,
|
||
model_revision: Optional[str] = None,
|
||
) -> Any:
|
||
"""创建 modelscope pipeline,并在需要时把底层模型迁移到目标设备。"""
|
||
from modelscope.pipelines import pipeline
|
||
|
||
pipeline_kwargs: dict[str, Any] = {
|
||
"task": task,
|
||
"model": model,
|
||
"device": modelscope_device,
|
||
}
|
||
if model_revision is not None:
|
||
pipeline_kwargs["model_revision"] = model_revision
|
||
|
||
pipeline_instance = pipeline(**pipeline_kwargs)
|
||
_move_pipeline_model_to_device(pipeline_instance, modelscope_device)
|
||
return pipeline_instance
|
||
|
||
|
||
def _enable_batched_sv(
|
||
pipeline_instance: Any,
|
||
modelscope_device: str,
|
||
max_batch_size: int = 32,
|
||
) -> Any:
|
||
"""
|
||
对说话人分离 pipeline 启用 batched SV 推理。
|
||
|
||
原始 pipeline 的 forward 方法逐个 segment 调用 sv_pipeline 提取 embedding,
|
||
这里改为将所有 segment 拼成一个 batch 一次性推理,大幅减少 GPU 调用次数。
|
||
同时将子 pipeline(sv / vad / change_locator)绑定到指定 device。
|
||
|
||
Args:
|
||
pipeline_instance: CAM++ diarization pipeline 实例
|
||
modelscope_device: 设备名称
|
||
max_batch_size: 最大批处理大小,防止 OOM
|
||
"""
|
||
if getattr(pipeline_instance, "_batched_sv_enabled", False):
|
||
return pipeline_instance
|
||
|
||
from modelscope.utils.constant import Tasks
|
||
|
||
config = getattr(pipeline_instance, "config", None)
|
||
if not isinstance(config, Mapping):
|
||
logger.warning("CAM++ pipeline 缺少可读取的 config,跳过 batched SV 优化")
|
||
return pipeline_instance
|
||
|
||
sv_model = config.get("speaker_model")
|
||
vad_model = config.get("vad_model")
|
||
change_locator = config.get("change_locator")
|
||
|
||
if isinstance(sv_model, str) and sv_model:
|
||
pipeline_instance.sv_pipeline = _create_modelscope_pipeline(
|
||
task=Tasks.speaker_verification,
|
||
model=sv_model,
|
||
modelscope_device=modelscope_device,
|
||
)
|
||
|
||
if isinstance(vad_model, str) and vad_model:
|
||
pipeline_instance.vad_pipeline = _create_modelscope_pipeline(
|
||
task=Tasks.voice_activity_detection,
|
||
model=vad_model,
|
||
modelscope_device=modelscope_device,
|
||
model_revision="v2.0.2",
|
||
)
|
||
|
||
if isinstance(change_locator, str) and change_locator:
|
||
pipeline_instance.change_locator_pipeline = _create_modelscope_pipeline(
|
||
task=Tasks.speaker_diarization,
|
||
model=change_locator,
|
||
modelscope_device=modelscope_device,
|
||
)
|
||
|
||
def batched_forward(self: Any, segments: Sequence[Sequence[Any]]) -> np.ndarray:
|
||
"""批量提取说话人 embedding,替代逐段串行推理"""
|
||
sv_model_instance = getattr(getattr(self, "sv_pipeline", None), "model", None)
|
||
emb_size = int(getattr(sv_model_instance, "emb_size", 192))
|
||
|
||
if not segments:
|
||
return np.empty((0, emb_size), dtype=np.float32)
|
||
|
||
if sv_model_instance is None:
|
||
raise RuntimeError("CAM++ sv_pipeline.model 未初始化")
|
||
|
||
all_embeddings: list[np.ndarray] = []
|
||
total_segments = len(segments)
|
||
start_idx = 0
|
||
|
||
while start_idx < total_segments:
|
||
end_idx = min(start_idx + max_batch_size, total_segments)
|
||
batch_segments = segments[start_idx:end_idx]
|
||
|
||
batch_items: list[np.ndarray] = []
|
||
for segment in batch_segments:
|
||
if len(segment) < 3:
|
||
continue
|
||
batch_items.append(np.asarray(segment[2], dtype=np.float32))
|
||
|
||
if not batch_items:
|
||
start_idx = end_idx
|
||
continue
|
||
|
||
batch = np.stack(batch_items, axis=0)
|
||
|
||
with torch.no_grad():
|
||
embeddings = sv_model_instance(
|
||
cast(Any, torch).as_tensor(batch).to(modelscope_device)
|
||
)
|
||
|
||
if isinstance(embeddings, torch.Tensor):
|
||
all_embeddings.append(embeddings.detach().cpu().numpy())
|
||
else:
|
||
all_embeddings.append(np.asarray(embeddings, dtype=np.float32))
|
||
|
||
start_idx = end_idx
|
||
|
||
if not all_embeddings:
|
||
return np.empty((0, emb_size), dtype=np.float32)
|
||
|
||
return (
|
||
np.concatenate(all_embeddings, axis=0)
|
||
if len(all_embeddings) > 1
|
||
else all_embeddings[0]
|
||
)
|
||
|
||
import types
|
||
|
||
pipeline_instance.forward = types.MethodType(batched_forward, pipeline_instance)
|
||
pipeline_instance._batched_sv_enabled = True
|
||
|
||
logger.info(
|
||
"CAM++ 说话人分离启用 batched SV: device={}, sv_device={}, vad_device={}",
|
||
modelscope_device,
|
||
getattr(getattr(pipeline_instance, "sv_pipeline", None), "device_name", "unknown"),
|
||
getattr(getattr(pipeline_instance, "vad_pipeline", None), "device_name", "unknown"),
|
||
)
|
||
return pipeline_instance
|
||
|
||
|
||
def get_global_diarization_pipeline() -> Any:
|
||
"""获取全局说话人分离 pipeline(懒加载单例)"""
|
||
global _global_diarization_pipeline
|
||
|
||
with _diarization_pipeline_lock:
|
||
if _global_diarization_pipeline is None:
|
||
try:
|
||
from modelscope.utils.constant import Tasks
|
||
from ..infrastructure.model_utils import resolve_model_path
|
||
|
||
model_id = 'iic/speech_campplus_speaker-diarization_common'
|
||
model_path = resolve_model_path(model_id)
|
||
modelscope_device = _resolve_modelscope_device()
|
||
|
||
logger.info(
|
||
"正在加载 CAM++ 说话人分离模型: {}, device={}",
|
||
model_path,
|
||
modelscope_device,
|
||
)
|
||
_global_diarization_pipeline = _create_modelscope_pipeline(
|
||
task=Tasks.speaker_diarization,
|
||
model=model_path,
|
||
modelscope_device=modelscope_device,
|
||
)
|
||
_global_diarization_pipeline = _enable_batched_sv(
|
||
_global_diarization_pipeline, modelscope_device
|
||
)
|
||
logger.info("CAM++ 模型加载成功(已启用 batched SV)")
|
||
except Exception as e:
|
||
logger.error(f"CAM++ 模型加载失败: {e}")
|
||
raise DefaultServerErrorException(f"说话人分离模型加载失败: {str(e)}")
|
||
|
||
return _global_diarization_pipeline
|
||
|
||
|
||
class SpeakerDiarizer:
|
||
"""基于 CAM++ 的说话人分离器"""
|
||
|
||
DEFAULT_MIN_SEGMENT_SEC = 1.0
|
||
DEFAULT_SAMPLE_RATE = 16000
|
||
LOW_ENERGY_SEARCH_WINDOW_MS = 10000
|
||
LOW_ENERGY_CONTEXT_MS = 160
|
||
LOW_ENERGY_STEP_MS = 20
|
||
|
||
def __init__(
|
||
self,
|
||
min_segment_sec: float = DEFAULT_MIN_SEGMENT_SEC,
|
||
):
|
||
self.min_segment_sec = min_segment_sec
|
||
self.min_segment_ms = int(min_segment_sec * 1000)
|
||
|
||
def diarize(
|
||
self, audio_path: str
|
||
) -> List[SpeakerSegment]:
|
||
"""执行说话人分离
|
||
|
||
Args:
|
||
audio_path: 音频文件路径
|
||
|
||
Returns:
|
||
原始分段列表(未合并)
|
||
"""
|
||
audio_duration_ms: Optional[int] = None
|
||
try:
|
||
try:
|
||
audio_duration_ms = int(librosa.get_duration(path=audio_path) * 1000)
|
||
except Exception:
|
||
audio_duration_ms = None
|
||
|
||
# CAM++ 对极短片段收益很低,且容易直接报 "too short"。
|
||
# 这里提前降级成单说话人,避免无意义 warning 刷屏。
|
||
if audio_duration_ms is not None and audio_duration_ms < self.min_segment_ms:
|
||
logger.debug(
|
||
"音频时长过短,跳过 CAM++ 说话人分离: duration_ms=%s < min_segment_ms=%s",
|
||
audio_duration_ms,
|
||
self.min_segment_ms,
|
||
)
|
||
return [
|
||
SpeakerSegment(
|
||
start_ms=0,
|
||
end_ms=max(audio_duration_ms, 1),
|
||
speaker_id="说话人1",
|
||
)
|
||
]
|
||
|
||
pipeline = get_global_diarization_pipeline()
|
||
|
||
logger.info(f"开始说话人分离: {audio_path}")
|
||
with _diarization_inference_semaphore:
|
||
result = pipeline(audio_path)
|
||
|
||
# 解析结果: {'text': [[start, end, speaker_id], ...]}
|
||
# pipeline 返回类型不确定,需要安全地获取 'text' 字段
|
||
if isinstance(result, dict):
|
||
raw_output = result.get('text', [])
|
||
else:
|
||
raw_output = getattr(result, 'text', []) or []
|
||
|
||
segments = []
|
||
for seg in raw_output:
|
||
if isinstance(seg, list) and len(seg) == 3:
|
||
try:
|
||
start_ms = int(float(seg[0]) * 1000)
|
||
end_ms = int(float(seg[1]) * 1000)
|
||
speaker_id = f"说话人{int(seg[2]) + 1}"
|
||
segments.append(SpeakerSegment(
|
||
start_ms=start_ms,
|
||
end_ms=end_ms,
|
||
speaker_id=speaker_id,
|
||
))
|
||
except (ValueError, TypeError) as e:
|
||
logger.warning(f"跳过格式错误的片段: {seg}, 错误: {e}")
|
||
|
||
logger.info(f"说话人分离完成,原始片段数: {len(segments)}")
|
||
# 诊断日志:打印前20个原始片段
|
||
for i, seg in enumerate(segments[:20]):
|
||
logger.debug(
|
||
f"[CAM++原始] #{i}: {seg.start_sec:.2f}-{seg.end_sec:.2f}s "
|
||
f"({seg.duration_sec:.2f}s) {seg.speaker_id}"
|
||
)
|
||
return segments
|
||
|
||
except Exception as e:
|
||
error_msg = str(e).lower()
|
||
|
||
# 音频太短时,返回默认的单说话人片段
|
||
if "too short" in error_msg:
|
||
logger.debug("CAM++ 跳过过短音频,回退单说话人片段: %s", e)
|
||
if audio_duration_ms is None:
|
||
try:
|
||
audio_duration_ms = int(librosa.get_duration(path=audio_path) * 1000)
|
||
except Exception:
|
||
audio_duration_ms = 5000
|
||
|
||
return [
|
||
SpeakerSegment(
|
||
start_ms=0,
|
||
end_ms=audio_duration_ms,
|
||
speaker_id="说话人1",
|
||
)
|
||
]
|
||
|
||
# 其他异常正常抛出
|
||
logger.error(f"说话人分离失败: {e}")
|
||
raise DefaultServerErrorException(f"说话人分离失败: {str(e)}")
|
||
|
||
def merge_consecutive_segments(
|
||
self, segments: List[SpeakerSegment]
|
||
) -> List[SpeakerSegment]:
|
||
"""合并同一说话人的连续片段"""
|
||
if not segments:
|
||
return []
|
||
|
||
# 按开始时间排序
|
||
sorted_segments = sorted(segments, key=lambda x: x.start_ms)
|
||
|
||
merged = []
|
||
current = SpeakerSegment(
|
||
start_ms=sorted_segments[0].start_ms,
|
||
end_ms=sorted_segments[0].end_ms,
|
||
speaker_id=sorted_segments[0].speaker_id,
|
||
)
|
||
|
||
for seg in sorted_segments[1:]:
|
||
if seg.speaker_id == current.speaker_id:
|
||
# 同一说话人,扩展结束时间
|
||
current.end_ms = max(current.end_ms, seg.end_ms)
|
||
else:
|
||
# 不同说话人,保存当前段,开始新段
|
||
logger.debug(
|
||
f"[合并中断] 说话人切换: {current.speaker_id} → {seg.speaker_id} "
|
||
f"在 {seg.start_sec:.2f}s,保存片段 {current.start_sec:.2f}-{current.end_sec:.2f}s"
|
||
)
|
||
merged.append(current)
|
||
current = SpeakerSegment(
|
||
start_ms=seg.start_ms,
|
||
end_ms=seg.end_ms,
|
||
speaker_id=seg.speaker_id,
|
||
)
|
||
|
||
# 保存最后一段
|
||
merged.append(current)
|
||
|
||
logger.info(f"合并同一说话人连续片段: {len(segments)} → {len(merged)}")
|
||
# 诊断日志:打印合并后的前20个片段
|
||
for i, seg in enumerate(merged[:20]):
|
||
logger.debug(
|
||
f"[合并后] #{i}: {seg.start_sec:.2f}-{seg.end_sec:.2f}s "
|
||
f"({seg.duration_sec:.2f}s) {seg.speaker_id}"
|
||
)
|
||
return merged
|
||
|
||
def merge_short_segments(
|
||
self, segments: List[SpeakerSegment]
|
||
) -> List[SpeakerSegment]:
|
||
"""智能合并短片段
|
||
|
||
策略:
|
||
1. 第一层:<10s的片段向后合并(避免孤立短片段)
|
||
2. 第二层:60s累积合并(合并连续片段)
|
||
"""
|
||
if not segments:
|
||
return []
|
||
|
||
max_segment_sec = settings.MAX_SEGMENT_SEC
|
||
|
||
# 按开始时间排序
|
||
sorted_segments = sorted(segments, key=lambda x: x.start_ms)
|
||
|
||
# 第一层:<10s累积向后合并(循环计算直到>=10s或超过60s)
|
||
merged = []
|
||
i = 0
|
||
while i < len(sorted_segments):
|
||
seg = sorted_segments[i]
|
||
|
||
# 如果>=10s,直接添加
|
||
if seg.duration_sec >= 10.0:
|
||
merged.append(seg)
|
||
i += 1
|
||
continue
|
||
|
||
# <10s,开始累积合并
|
||
current_start_ms = seg.start_ms
|
||
current_end_ms = seg.end_ms
|
||
current_duration_sec = seg.duration_sec
|
||
j = i + 1
|
||
|
||
# 累积合并,只要<10s且同说话人且不超过60s
|
||
while j < len(sorted_segments) and current_duration_sec < 10.0:
|
||
next_seg = sorted_segments[j]
|
||
if next_seg.speaker_id != seg.speaker_id:
|
||
break
|
||
new_duration = (next_seg.end_ms - current_start_ms) / 1000.0
|
||
if new_duration > max_segment_sec:
|
||
break
|
||
current_end_ms = next_seg.end_ms
|
||
current_duration_sec = new_duration
|
||
j += 1
|
||
|
||
# 创建合并后的片段
|
||
merged_seg = SpeakerSegment(
|
||
start_ms=current_start_ms,
|
||
end_ms=current_end_ms,
|
||
speaker_id=seg.speaker_id,
|
||
)
|
||
merged.append(merged_seg)
|
||
|
||
if j > i + 1:
|
||
logger.debug(
|
||
f"[第一层] {seg.speaker_id}: "
|
||
f"累积合并了 {j - i} 个片段,结果 {merged_seg.duration_sec:.1f}s"
|
||
)
|
||
i = j
|
||
|
||
# 第二层:60s累积合并
|
||
final_merged = []
|
||
i = 0
|
||
while i < len(merged):
|
||
seg = merged[i]
|
||
current_start_ms = seg.start_ms
|
||
current_end_ms = seg.end_ms
|
||
j = i + 1
|
||
|
||
# 累积合并,只要 <= 60s 且同说话人
|
||
while j < len(merged):
|
||
next_seg = merged[j]
|
||
if next_seg.speaker_id != seg.speaker_id:
|
||
break
|
||
new_duration = (next_seg.end_ms - current_start_ms) / 1000.0
|
||
if new_duration > max_segment_sec:
|
||
break
|
||
current_end_ms = next_seg.end_ms
|
||
j += 1
|
||
|
||
merged_seg = SpeakerSegment(
|
||
start_ms=current_start_ms,
|
||
end_ms=current_end_ms,
|
||
speaker_id=seg.speaker_id,
|
||
)
|
||
final_merged.append(merged_seg)
|
||
|
||
if j > i + 1:
|
||
logger.debug(
|
||
f"[第二层] {seg.speaker_id}: "
|
||
f"合并了 {j - i} 个片段"
|
||
)
|
||
i = j
|
||
|
||
return final_merged
|
||
|
||
def _find_low_energy_boundary_ms(
|
||
self,
|
||
audio_data: np.ndarray,
|
||
sample_rate: int,
|
||
lower_ms: int,
|
||
upper_ms: int,
|
||
) -> int:
|
||
lower_ms = max(0, lower_ms)
|
||
upper_ms = max(lower_ms, upper_ms)
|
||
if sample_rate <= 0 or audio_data.size == 0:
|
||
return upper_ms
|
||
|
||
context_samples = max(
|
||
1, int(sample_rate * self.LOW_ENERGY_CONTEXT_MS / 1000)
|
||
)
|
||
candidate_points = list(
|
||
range(lower_ms, upper_ms + 1, self.LOW_ENERGY_STEP_MS)
|
||
)
|
||
if not candidate_points or candidate_points[-1] != upper_ms:
|
||
candidate_points.append(upper_ms)
|
||
|
||
best_ms = upper_ms
|
||
best_energy = float("inf")
|
||
audio_length = int(audio_data.shape[0])
|
||
|
||
for candidate_ms in candidate_points:
|
||
center_sample = int(candidate_ms * sample_rate / 1000)
|
||
start_sample = max(0, center_sample - context_samples // 2)
|
||
end_sample = min(audio_length, center_sample + context_samples // 2)
|
||
if start_sample >= end_sample:
|
||
continue
|
||
|
||
window = audio_data[start_sample:end_sample]
|
||
energy = float(np.mean(np.square(window)))
|
||
if energy <= best_energy:
|
||
best_energy = energy
|
||
best_ms = candidate_ms
|
||
|
||
return best_ms
|
||
|
||
def split_long_segments(
|
||
self,
|
||
segments: List[SpeakerSegment],
|
||
audio_data: np.ndarray,
|
||
sample_rate: int,
|
||
) -> List[SpeakerSegment]:
|
||
max_segment_ms = int(settings.MAX_SEGMENT_SEC * 1000)
|
||
if max_segment_ms <= 0:
|
||
return segments
|
||
|
||
split_segments: List[SpeakerSegment] = []
|
||
for seg in segments:
|
||
if seg.duration_ms <= max_segment_ms:
|
||
split_segments.append(seg)
|
||
continue
|
||
|
||
current_start_ms = seg.start_ms
|
||
while seg.end_ms - current_start_ms > max_segment_ms:
|
||
hard_boundary_ms = current_start_ms + max_segment_ms
|
||
lower_boundary_ms = max(
|
||
current_start_ms + self.min_segment_ms,
|
||
hard_boundary_ms - self.LOW_ENERGY_SEARCH_WINDOW_MS,
|
||
)
|
||
boundary_ms = self._find_low_energy_boundary_ms(
|
||
audio_data=audio_data,
|
||
sample_rate=sample_rate,
|
||
lower_ms=lower_boundary_ms,
|
||
upper_ms=hard_boundary_ms,
|
||
)
|
||
if boundary_ms <= current_start_ms:
|
||
boundary_ms = hard_boundary_ms
|
||
|
||
split_segments.append(
|
||
SpeakerSegment(
|
||
start_ms=current_start_ms,
|
||
end_ms=boundary_ms,
|
||
speaker_id=seg.speaker_id,
|
||
)
|
||
)
|
||
current_start_ms = boundary_ms
|
||
|
||
remaining_ms = seg.end_ms - current_start_ms
|
||
if remaining_ms >= self.min_segment_ms:
|
||
split_segments.append(
|
||
SpeakerSegment(
|
||
start_ms=current_start_ms,
|
||
end_ms=seg.end_ms,
|
||
speaker_id=seg.speaker_id,
|
||
)
|
||
)
|
||
elif split_segments:
|
||
split_segments[-1].end_ms = seg.end_ms
|
||
|
||
if len(split_segments) != len(segments):
|
||
logger.info(
|
||
"Split long speaker segments by low energy: {} -> {}, max={}s",
|
||
len(segments),
|
||
len(split_segments),
|
||
settings.MAX_SEGMENT_SEC,
|
||
)
|
||
return split_segments
|
||
|
||
def split_audio_by_speakers(
|
||
self,
|
||
audio_path: str,
|
||
output_dir: Optional[str] = None,
|
||
) -> List[SpeakerSegment]:
|
||
"""完整的说话人分离流程
|
||
|
||
流程:
|
||
1. 执行CAM++说话人分离
|
||
2. 智能合并短片段(两层合并策略)
|
||
- 第一层:<10s片段累积合并
|
||
- 第二层:60s累积合并
|
||
3. 提取音频数据,保存临时文件
|
||
|
||
Args:
|
||
audio_path: 音频文件路径
|
||
output_dir: 输出目录
|
||
|
||
Returns:
|
||
SpeakerSegment 列表
|
||
"""
|
||
try:
|
||
# 1. 执行说话人分离
|
||
raw_segments = self.diarize(audio_path)
|
||
|
||
if not raw_segments:
|
||
logger.warning("说话人分离未检测到任何片段")
|
||
return []
|
||
|
||
# 2. 智能合并短片段(第一个<10s的同说话人片段向后合并)
|
||
final_segments = self.merge_short_segments(raw_segments)
|
||
|
||
# 3. Load audio before low-energy splitting and segment extraction.
|
||
logger.info("加载音频并提取片段...")
|
||
audio_data, sr = librosa.load(audio_path, sr=self.DEFAULT_SAMPLE_RATE)
|
||
sample_rate = int(sr)
|
||
|
||
final_segments = self.split_long_segments(
|
||
final_segments,
|
||
audio_data,
|
||
sample_rate,
|
||
)
|
||
logger.info(f"智能合并完成: {len(raw_segments)} → {len(final_segments)} 个片段")
|
||
|
||
output_dir = output_dir or settings.TEMP_DIR
|
||
os.makedirs(output_dir, exist_ok=True)
|
||
|
||
for idx, seg in enumerate(final_segments):
|
||
start_sample = int(seg.start_ms / 1000 * sample_rate)
|
||
end_sample = int(seg.end_ms / 1000 * sample_rate)
|
||
|
||
seg.audio_data = audio_data[start_sample:end_sample]
|
||
|
||
# 保存临时文件
|
||
temp_file = tempfile.NamedTemporaryFile(
|
||
delete=False,
|
||
suffix=".wav",
|
||
dir=output_dir,
|
||
prefix=f"{seg.speaker_id}_{idx:03d}_",
|
||
)
|
||
temp_path = temp_file.name
|
||
temp_file.close()
|
||
|
||
sf.write(temp_path, seg.audio_data, sample_rate)
|
||
seg.temp_file = temp_path
|
||
|
||
# 统计
|
||
unique_speakers = sorted(set(seg.speaker_id for seg in final_segments))
|
||
logger.info(
|
||
f"音频分割完成: {len(final_segments)} 个片段, "
|
||
f"{len(unique_speakers)} 个说话人"
|
||
)
|
||
for spk in unique_speakers:
|
||
spk_segs = [s for s in final_segments if s.speaker_id == spk]
|
||
total_time = sum(s.duration_sec for s in spk_segs)
|
||
logger.info(f" {spk}: {len(spk_segs)} 片段, {total_time:.2f}s")
|
||
|
||
return final_segments
|
||
|
||
except Exception as e:
|
||
logger.error(f"说话人分离流程失败: {e}")
|
||
raise DefaultServerErrorException(f"说话人分离失败: {str(e)}")
|
||
|
||
@staticmethod
|
||
def cleanup_segments(segments: List[SpeakerSegment]) -> None:
|
||
"""清理临时文件"""
|
||
for seg in segments:
|
||
if seg.temp_file and os.path.exists(seg.temp_file):
|
||
try:
|
||
os.remove(seg.temp_file)
|
||
except Exception as e:
|
||
logger.warning(f"清理临时文件失败: {seg.temp_file}, {e}")
|