/
/
/
1"""
2Vocal activity inference using FireRed AED.
3
4The model architecture and weights are adapted from FireRedVAD:
5https://github.com/FireRedTeam/FireRedVAD
6Copyright 2026 Xiaohongshu. Licensed under Apache-2.0.
7"""
8
9from __future__ import annotations
10
11import math
12from collections.abc import Iterator
13from pathlib import Path
14
15import kaldi_native_fbank as knf
16import numpy as np
17import numpy.typing as npt
18import torch
19from torch import nn
20from torch.nn import functional
21
22FIRERED_SAMPLE_RATE = 16000
23FIRERED_FRAME_DURATION = 0.1
24FIRERED_MEL_BINS = 80
25FIRERED_OUTPUT_CLASSES = 3
26FIRERED_PARAMETER_COUNT = 588_931
27FIRERED_MAX_INFERENCE_FRAMES = 30_000
28# Covers the eight-layer FSMN receptive field in both directions.
29FIRERED_INFERENCE_CONTEXT_FRAMES = 160
30
31_MODEL_PATH = Path(__file__).parent / "resources" / "firered_aed.pt"
32_CMVN_PATH = Path(__file__).parent / "resources" / "firered_aed_cmvn.npz"
33
34
35class FireRedFbank:
36 """Streaming FireRed AED feature extractor for 16 kHz mono PCM."""
37
38 def __init__(
39 self,
40 cmvn_means: npt.NDArray[np.float64],
41 cmvn_inverse_std: npt.NDArray[np.float64],
42 ) -> None:
43 """
44 Initialize the feature extractor.
45
46 :param cmvn_means: Fixed FireRed feature means.
47 :param cmvn_inverse_std: Fixed FireRed inverse standard deviations.
48 """
49 if cmvn_means.shape != (FIRERED_MEL_BINS,) or cmvn_inverse_std.shape != (FIRERED_MEL_BINS,):
50 raise ValueError("invalid FireRed CMVN dimensions")
51 options = knf.FbankOptions()
52 options.frame_opts.samp_freq = FIRERED_SAMPLE_RATE
53 options.frame_opts.frame_length_ms = 25
54 options.frame_opts.frame_shift_ms = 10
55 options.frame_opts.dither = 0
56 options.frame_opts.snip_edges = True
57 options.mel_opts.num_bins = FIRERED_MEL_BINS
58 options.mel_opts.debug_mel = False
59 self._fbank = knf.OnlineFbank(options)
60 self._cmvn_means = cmvn_means
61 self._cmvn_inverse_std = cmvn_inverse_std
62 self._next_frame = 0
63 self._finished = False
64
65 def process(self, pcm_16k: npt.NDArray[np.float32]) -> npt.NDArray[np.float32]:
66 """
67 Consume normalized 16 kHz mono PCM and return newly available features.
68
69 :param pcm_16k: Normalized mono float32 samples at 16 kHz.
70 """
71 if self._finished:
72 raise RuntimeError("FireRed feature extraction is already finalized")
73 if pcm_16k.size:
74 quantized = quantize_pcm_for_firered(pcm_16k)
75 self._fbank.accept_waveform(FIRERED_SAMPLE_RATE, quantized.tolist())
76 return self._collect_ready_frames()
77
78 def finalize(self) -> npt.NDArray[np.float32]:
79 """Finish the stream and return any final features."""
80 if not self._finished:
81 self._fbank.input_finished()
82 self._finished = True
83 return self._collect_ready_frames()
84
85 def _collect_ready_frames(self) -> npt.NDArray[np.float32]:
86 ready = self._fbank.num_frames_ready
87 if ready <= self._next_frame:
88 return np.empty((0, FIRERED_MEL_BINS), dtype=np.float32)
89 features = np.vstack(
90 [self._fbank.get_frame(index) for index in range(self._next_frame, ready)]
91 )
92 self._fbank.pop(ready - self._next_frame)
93 self._next_frame = ready
94 normalized = (features.astype(np.float64) - self._cmvn_means) * self._cmvn_inverse_std
95 return normalized.astype(np.float32)
96
97
98def load_firered_components(
99 device: str = "cpu",
100) -> tuple[nn.Module, npt.NDArray[np.float64], npt.NDArray[np.float64]]:
101 """
102 Load the FireRed AED model and fixed CMVN statistics.
103
104 :param device: Device used for model inference.
105 """
106 model = _FireRedAedModel().to(device)
107 state_dict = torch.load(_MODEL_PATH, map_location=device, weights_only=True)
108 model.load_state_dict(state_dict, strict=True)
109 model.eval()
110 parameter_count = sum(parameter.numel() for parameter in model.parameters())
111 if parameter_count != FIRERED_PARAMETER_COUNT:
112 raise RuntimeError(
113 f"invalid FireRed AED parameter count: {parameter_count} "
114 f"(expected {FIRERED_PARAMETER_COUNT})"
115 )
116 with np.load(_CMVN_PATH) as cmvn:
117 means = cmvn["means"].astype(np.float64)
118 inverse_std = cmvn["inverse_std"].astype(np.float64)
119 if means.shape != (FIRERED_MEL_BINS,) or inverse_std.shape != (FIRERED_MEL_BINS,):
120 raise RuntimeError("invalid FireRed AED CMVN resource")
121 return model, means, inverse_std
122
123
124def quantize_pcm_for_firered(
125 pcm: npt.NDArray[np.float32],
126) -> npt.NDArray[np.float32]:
127 """
128 Convert normalized float PCM to the signed 16-bit scale expected by FireRed.
129
130 :param pcm: Normalized mono float PCM.
131 """
132 scaled = np.rint(pcm.astype(np.float64) * 32768.0)
133 np.clip(scaled, -32768.0, 32767.0, out=scaled)
134 return scaled.astype(np.float32)
135
136
137def split_firered_features(
138 features: npt.NDArray[np.float32],
139) -> Iterator[tuple[npt.NDArray[np.float32], int, int]]:
140 """
141 Yield inference chunks with context and the core output range.
142
143 :param features: CMVN-normalized FireRed features shaped ``(frames, 80)``.
144 """
145 total_frames = len(features)
146 for core_start in range(0, total_frames, FIRERED_MAX_INFERENCE_FRAMES):
147 core_end = min(core_start + FIRERED_MAX_INFERENCE_FRAMES, total_frames)
148 chunk_start = max(0, core_start - FIRERED_INFERENCE_CONTEXT_FRAMES)
149 chunk_end = min(total_frames, core_end + FIRERED_INFERENCE_CONTEXT_FRAMES)
150 yield features[chunk_start:chunk_end], core_start - chunk_start, core_end - core_start
151
152
153def infer_firered_chunk(
154 model: nn.Module,
155 features: npt.NDArray[np.float32],
156 device: str = "cpu",
157) -> npt.NDArray[np.float32]:
158 """
159 Run one FireRed AED model call.
160
161 :param model: Loaded FireRed AED model.
162 :param features: Feature chunk shaped ``(frames, 80)``.
163 :param device: Device used for model inference.
164 """
165 if features.ndim != 2 or features.shape[1] != FIRERED_MEL_BINS:
166 raise ValueError("invalid FireRed feature shape")
167 if not len(features):
168 return np.empty((0, FIRERED_OUTPUT_CLASSES), dtype=np.float32)
169 tensor = torch.from_numpy(features).unsqueeze(0).to(device)
170 with torch.inference_mode():
171 probabilities = model(tensor)
172 result: npt.NDArray[np.float32] = np.asarray(
173 probabilities.squeeze(0).cpu().numpy(),
174 dtype=np.float32,
175 )
176 if result.shape != (len(features), FIRERED_OUTPUT_CLASSES):
177 raise RuntimeError("invalid FireRed AED output shape")
178 return result
179
180
181def vocal_activity_probabilities(
182 frame_probabilities: npt.NDArray[np.float32],
183 duration: float,
184) -> npt.NDArray[np.float32]:
185 """
186 Convert 10 ms FireRed outputs to a 100 ms vocal timeline.
187
188 :param frame_probabilities: Speech, singing, and music probabilities per fbank frame.
189 :param duration: Canonical analysis duration in seconds.
190 """
191 if not math.isfinite(duration) or duration < 0:
192 raise ValueError("invalid vocal activity duration")
193 output_frames = math.ceil(duration / FIRERED_FRAME_DURATION)
194 if output_frames == 0:
195 return np.empty(0, dtype=np.float32)
196 if frame_probabilities.ndim != 2 or frame_probabilities.shape[1] != (FIRERED_OUTPUT_CLASSES):
197 raise ValueError("invalid FireRed AED probability shape")
198 if not np.isfinite(frame_probabilities).all():
199 raise ValueError("FireRed AED produced non-finite probabilities")
200 if not len(frame_probabilities):
201 return np.zeros(output_frames, dtype=np.float32)
202
203 clipped = np.clip(frame_probabilities, 0.0, 1.0)
204 vocal_frames = np.maximum(clipped[:, 0], clipped[:, 1])
205 starts = np.arange(0, len(vocal_frames), 10)
206 sums = np.add.reduceat(vocal_frames.astype(np.float64), starts)
207 counts = np.minimum(10, len(vocal_frames) - starts)
208 probabilities = (sums / counts).astype(np.float32)
209 if len(probabilities) < output_frames:
210 probabilities = np.pad(
211 probabilities,
212 (0, output_frames - len(probabilities)),
213 mode="edge",
214 )
215 else:
216 probabilities = probabilities[:output_frames]
217 if not np.isfinite(probabilities).all():
218 raise ValueError("invalid vocal activity probabilities")
219 result: npt.NDArray[np.float32] = np.asarray(
220 np.clip(probabilities, 0.0, 1.0),
221 dtype=np.float32,
222 )
223 return result
224
225
226class _FireRedAedModel(nn.Module):
227 """FireRed AED detection model."""
228
229 def __init__(self) -> None:
230 super().__init__()
231 self.dfsmn = _Dfsmn()
232 self.out = nn.Linear(256, FIRERED_OUTPUT_CLASSES)
233
234 def forward(self, features: torch.Tensor) -> torch.Tensor:
235 """Return speech, singing, and music probabilities."""
236 hidden: torch.Tensor = self.dfsmn(features)
237 logits: torch.Tensor = self.out(hidden)
238 return torch.sigmoid(logits)
239
240
241class _Dfsmn(nn.Module):
242 """Deep feedforward sequential memory network used by FireRed AED."""
243
244 def __init__(self) -> None:
245 super().__init__()
246 self.fc1 = nn.Sequential(nn.Linear(FIRERED_MEL_BINS, 256), nn.ReLU(), nn.Dropout(0.05))
247 self.fc2 = nn.Sequential(nn.Linear(256, 128), nn.ReLU(), nn.Dropout(0.05))
248 self.fsmn1 = _Fsmn()
249 self.fsmns = nn.ModuleList([_DfsmnBlock() for _ in range(7)])
250 self.dnns = nn.Sequential(nn.Linear(128, 256), nn.ReLU(), nn.Dropout(0.05))
251
252 def forward(self, features: torch.Tensor) -> torch.Tensor:
253 memory = self.fsmn1(self.fc2(self.fc1(features)))
254 for block in self.fsmns:
255 memory = block(memory)
256 output: torch.Tensor = self.dnns(memory)
257 return output
258
259
260class _DfsmnBlock(nn.Module):
261 """Residual DFSMN block."""
262
263 def __init__(self) -> None:
264 super().__init__()
265 self.fc1 = nn.Sequential(nn.Linear(128, 256), nn.ReLU(), nn.Dropout(0.05))
266 self.fc2 = nn.Linear(256, 128, bias=False)
267 self.fsmn = _Fsmn()
268
269 def forward(self, features: torch.Tensor) -> torch.Tensor:
270 hidden: torch.Tensor = self.fc1(features)
271 projection: torch.Tensor = self.fc2(hidden)
272 memory: torch.Tensor = self.fsmn(projection)
273 return features + memory
274
275
276class _Fsmn(nn.Module):
277 """Depthwise bidirectional memory block."""
278
279 def __init__(self) -> None:
280 super().__init__()
281 self.lookback_filter = nn.Conv1d(
282 128,
283 128,
284 kernel_size=20,
285 padding=19,
286 groups=128,
287 bias=False,
288 )
289 self.lookahead_filter = nn.Conv1d(
290 128,
291 128,
292 kernel_size=20,
293 padding=19,
294 groups=128,
295 bias=False,
296 )
297
298 def forward(self, features: torch.Tensor) -> torch.Tensor:
299 transposed = features.permute(0, 2, 1).contiguous()
300 lookback: torch.Tensor = self.lookback_filter(transposed)[:, :, :-19]
301 memory: torch.Tensor = transposed + lookback
302 if features.size(1) > 1:
303 lookahead: torch.Tensor = self.lookahead_filter(transposed)
304 memory += functional.pad(lookahead[:, :, 20:], (0, 1))
305 return memory.permute(0, 2, 1).contiguous()
306