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