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