File size: 6,990 Bytes
906ada2 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 | # 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))
|