HY-2012's picture
Upload folder using huggingface_hub
906ada2 verified
Raw
History Blame Contribute Delete
6.99 kB
# Copyright 2026 AXERA-TECH (authors: Magnetar)
#
# CAMPPlus speaker embedding inference SDK (AXera NPU).
#
# Mirrors python/utils/ax_cam_bin.py (AX_SpeakerEmbeddingInference) of
# 3D-Speaker-MT.axera:
# - 16 kHz mono audio
# - 80-dim kaldi fbank (25 ms / 10 ms, dither=0, snip_edges=true, povey)
# - mean_nor: subtract per-bin mean over frames
# - chunk 1.5 s / stride 0.75 s, circle-pad to 57900 samples -> 360 frames
# - campplus.axmodel: feature [1,360,80] float32 -> embedding [1,192]
import os
import numpy as np
import torch
import torchaudio.compliance.kaldi as Kaldi
try:
import axengine as axe
except ImportError:
axe = None
FRAMES = 360
FEAT_DIM = 80
SAMPLE_RATE = 16000
MIN_WAV_LEN = 57900 # 57900 samples -> exactly 360 fbank frames
EMBEDDING_DIM = 192
def load_wav(path: str, target_sr: int = SAMPLE_RATE):
"""Read wav (any format supported by the backend), resample to target_sr,
return mono float tensor [1, T]. Falls back to the ffmpeg backend when the
default backend is unavailable (e.g. board without torchcodec)."""
import torchaudio
try:
wav, fs = torchaudio.load(path)
except Exception:
wav, fs = torchaudio.load(path, backend="ffmpeg")
if fs != target_sr:
wav = torchaudio.functional.resample(wav, fs, target_sr)
if wav.shape[0] > 1:
wav = wav[0:1]
return wav
def circle_pad(x: torch.Tensor, target_len: int, dim: int = 0) -> torch.Tensor:
"""Mirrors speakerlab.utils.utils.circle_pad: repeat until target_len,
then truncate (no-op if already long enough)."""
xlen = x.shape[dim]
if xlen >= target_len:
return x
n = int(np.ceil(target_len / xlen))
xcat = torch.cat([x for _ in range(n)], dim=dim)
return torch.narrow(xcat, dim, 0, target_len)
class FBank(object):
"""Mirrors speakerlab.process.processor.FBank(80, 16000, mean_nor=True)."""
def __init__(self, n_mels=FEAT_DIM, sample_rate=SAMPLE_RATE, mean_nor=True):
self.n_mels = n_mels
self.sample_rate = sample_rate
self.mean_nor = mean_nor
def __call__(self, wav: torch.Tensor, dither: int = 0) -> torch.Tensor:
assert self.sample_rate == SAMPLE_RATE
if len(wav.shape) == 1:
wav = wav.unsqueeze(0)
if wav.shape[0] > 1:
wav = wav[0, :].unsqueeze(0)
assert len(wav.shape) == 2 and wav.shape[0] == 1
feat = Kaldi.fbank(wav, num_mel_bins=self.n_mels,
sample_frequency=SAMPLE_RATE, dither=dither)
if self.mean_nor:
feat = feat - feat.mean(0, keepdim=True)
return feat # [T, 80]
def chunk(st, ed, dur=1.5, step=0.75):
"""Mirrors utils/ax_cam_bin.py chunk(): sliding windows in seconds."""
chunks = []
subseg_st = st
while subseg_st + dur < ed + step:
subseg_ed = min(subseg_st + dur, ed)
chunks.append([subseg_st, subseg_ed])
subseg_st += step
return chunks
class CampplusModel:
"""Speaker embedding model running on AXera NPU via axengine.
Mirrors AX_SpeakerEmbeddingInference:
model = CampplusModel("models") # loads models/campplus.axmodel
embeddings = model(speech, 16000, chunks=[[0.0, 1.5], ...])
"""
def __init__(self, model_dir: str, model_file: str = "campplus.axmodel"):
if axe is None:
raise RuntimeError(
"axengine is not available; run on the AXera board "
"(pip install axengine)")
model_path = os.path.join(model_dir, model_file)
self.session = axe.InferenceSession(
model_path, providers="AxEngineExecutionProvider")
def infer(self, feats: np.ndarray) -> np.ndarray:
"""feats [B, 360, 80] float32 -> embedding [B, 192]."""
inputs = {self.session.get_inputs()[0].name:
np.ascontiguousarray(feats, dtype=np.float32)}
return self.session.run(None, inputs)[0]
def extract(self, wav: torch.Tensor) -> np.ndarray:
"""wav [1, T] or [T] (single chunk) -> embedding [1, 192].
Circle-pads to 57900 samples (360 fbank frames) when shorter.
"""
if len(wav.shape) == 1:
wav = wav.unsqueeze(0)
if wav.shape[0] > 1:
wav = wav[0:1] # mono
wav = circle_pad(wav[0], MIN_WAV_LEN).unsqueeze(0)
feature_extractor = FBank(FEAT_DIM, SAMPLE_RATE, mean_nor=True)
feats = torch.vmap(feature_extractor)(wav.unsqueeze(1))
if feats.shape[1] >= FRAMES:
feats = feats.narrow(1, 0, FRAMES)
else:
target_shape = list(feats.shape)
target_shape[1] = FRAMES
feats = feats.new_full(target_shape, 0.0)
return self.infer(feats.numpy())
def __call__(self, wav, fs: int = SAMPLE_RATE,
chunks=None, **kwargs) -> np.ndarray:
"""Extract speaker embeddings for each chunk.
Args:
wav: np.ndarray [T] mono audio, or path to a wav file
fs: sample rate (must be 16000)
chunks: list of [start_time, end_time] in seconds; if None the
whole audio is treated as one chunk
Returns:
embeddings np.ndarray [N, 192]
"""
if isinstance(wav, str):
wav, fs = load_wav(wav)
if fs != SAMPLE_RATE:
raise ValueError(f"input sample rate {fs} != {SAMPLE_RATE}")
wav = wav.numpy()
wav = torch.from_numpy(wav)
if len(wav.shape) == 1:
wav = wav.unsqueeze(0)
if wav.shape[0] > 1:
wav = wav[0:1] # mono
if chunks is None:
chunks = [[0.0, wav.shape[1] / fs]]
wavs = [wav[0, int(st * fs):int(ed * fs)] for st, ed in chunks]
# Pad all chunks to the same length (>= 57900 -> 360 frames)
max_len = max([x.shape[0] for x in wavs])
max_len = max(max_len, MIN_WAV_LEN)
wavs = [circle_pad(x, max_len) for x in wavs]
wavs = torch.stack(wavs).unsqueeze(1)
batch_size = 1 # onnx batch=1
embeddings = []
feature_extractor = FBank(FEAT_DIM, SAMPLE_RATE, mean_nor=True)
for i in range(0, len(wavs), batch_size):
batch_wavs = wavs[i:i + batch_size]
feats_batch = torch.vmap(feature_extractor)(batch_wavs)
if feats_batch.shape[1] >= FRAMES:
feats_batch = feats_batch.narrow(1, 0, FRAMES)
else:
target_shape = list(feats_batch.shape)
target_shape[1] = FRAMES
feats_batch = feats_batch.new_full(target_shape, 0.0)
embeddings.append(self.infer(feats_batch.numpy()))
return np.concatenate(embeddings, axis=0)
def cosine_similarity(a: np.ndarray, b: np.ndarray) -> float:
a = a.reshape(-1).astype(np.float32)
b = b.reshape(-1).astype(np.float32)
return float(np.dot(a, b) / (np.linalg.norm(a) * np.linalg.norm(b) + 1e-12))