Download python/campplus_sdk/inference.py from AXERA-TECH/campplus.AXERA: direct link, hf CLI and curl.
- Browser
- Download file 6.99 kB
-
https://huggingface.co/AXERA-TECH/campplus.AXERA/resolve/main/python/campplus_sdk/inference.py
- Command line
-
hf download hf://AXERA-TECH/campplus.AXERA/python/campplus_sdk/inference.py
-
curl -L -o inference.py https://huggingface.co/AXERA-TECH/campplus.AXERA/resolve/main/python/campplus_sdk/inference.py
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)) | |