# 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))