jgalego commited on
Commit
33f8f33
·
verified ·
1 Parent(s): b0ea6a4

Add MTRCNN-DG weights and evaluation

Browse files
Files changed (3) hide show
  1. model.pt +3 -0
  2. mtrcnn.py +1537 -0
  3. results/eval.json +1177 -0
model.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:1eeda46ea42e0dfe59f2724fc204728bfef7fe38ae3c3c5800cf00a2aa13712e
3
+ size 918678
mtrcnn.py ADDED
@@ -0,0 +1,1537 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # /// script
2
+ # requires-python = ">=3.12"
3
+ # dependencies = [
4
+ # "huggingface-hub",
5
+ # "jinja2",
6
+ # "numpy",
7
+ # "scipy",
8
+ # "soundfile",
9
+ # "torch",
10
+ # ]
11
+ # ///
12
+ """Train a compact domain-generalized mosquito species classifier."""
13
+
14
+ # pylint: disable=too-few-public-methods,too-many-arguments,too-many-instance-attributes
15
+ # pylint: disable=too-many-lines,too-many-locals,too-many-positional-arguments
16
+
17
+ import argparse
18
+ import json
19
+ import math
20
+ import os
21
+ import random
22
+ import re
23
+ import shutil
24
+ import statistics
25
+ import urllib.request
26
+ import zipfile
27
+ from collections import Counter
28
+ from pathlib import Path
29
+
30
+ import numpy as np
31
+ import soundfile
32
+ import torch
33
+ from huggingface_hub import (
34
+ EvalResult,
35
+ HfApi,
36
+ ModelCard,
37
+ ModelCardData,
38
+ hf_hub_download,
39
+ snapshot_download,
40
+ )
41
+ from huggingface_hub.errors import RepositoryNotFoundError
42
+ from scipy.signal import resample_poly
43
+ from torch import nn
44
+ from torch.nn import functional
45
+ from torch.utils.data import DataLoader, Dataset, WeightedRandomSampler
46
+
47
+ DATASET = "aptemvs/mosquitoes-biodcase2026-task5"
48
+ REPO = "jgalego/mosquito-mtrcnn-dg"
49
+ ZENODO = "https://zenodo.org/api/records/20478577/files/Development_data.zip/content"
50
+ SAMPLE_RATE = 8_000
51
+ N_FFT = 512
52
+ HOP_LENGTH = 80
53
+ N_MELS = 64
54
+ # Mains hum and handling noise sit below 200 Hz; wingbeat fundamentals are higher.
55
+ MIN_HZ = 200
56
+ LEVEL_RMS = 0.05
57
+ # Every D5 clip lasts 0.625 s; longer crops would pad them with telltale silence.
58
+ SECONDS = 0.625
59
+ SPECIES = [
60
+ "Aedes aegypti",
61
+ "Aedes albopictus",
62
+ "Culex quinquefasciatus",
63
+ "Anopheles gambiae",
64
+ "Anopheles arabiensis",
65
+ "Anopheles dirus",
66
+ "Culex pipiens",
67
+ "Anopheles minimus",
68
+ "Anopheles stephensi",
69
+ ]
70
+ DOMAINS = ["D1", "D2", "D3", "D4", "D5"]
71
+ FILE_ID = re.compile(r"^S_(\d+)_D_(\d+)_(\d+)$")
72
+ SYNTHETIC_DOMAINS = [
73
+ [0, 3, 4],
74
+ [0, 2, 4],
75
+ [3, 4],
76
+ [4],
77
+ [4],
78
+ [0],
79
+ [0, 1, 4],
80
+ [0, 2, 3],
81
+ [0, 1, 2],
82
+ ]
83
+
84
+
85
+ def parse_file_id(file_id):
86
+ """Return zero-based species and domain indices from an official file ID."""
87
+ match = FILE_ID.fullmatch(file_id)
88
+ if not match:
89
+ raise ValueError(f"invalid BioDCASE file ID: {file_id}")
90
+ species, domain, _ = map(int, match.groups())
91
+ if not (1 <= species <= len(SPECIES) and 1 <= domain <= len(DOMAINS)):
92
+ raise ValueError(f"out-of-range BioDCASE file ID: {file_id}")
93
+ return species - 1, domain - 1
94
+
95
+
96
+ def held_domain_split(ids, fold, seed=42, singleton_fraction=0.1):
97
+ """Hold out a rare domain, or rows when only one domain exists, per species."""
98
+ groups = {}
99
+ for file_id in ids:
100
+ groups.setdefault(parse_file_id(file_id), []).append(file_id)
101
+ held_domains = {}
102
+ validation_ids = []
103
+ for species, species_name in enumerate(SPECIES):
104
+ available = [domain for domain in range(len(DOMAINS)) if groups.get((species, domain))]
105
+ candidates = sorted(available)
106
+ if len(candidates) == 1:
107
+ held_domains[species] = None
108
+ rows = groups[species, candidates[0]].copy()
109
+ random.Random(seed + species).shuffle(rows)
110
+ validation_size = min(
111
+ len(rows) - 1,
112
+ max(1, round(singleton_fraction * len(rows))),
113
+ )
114
+ validation_ids.extend(rows[:validation_size])
115
+ continue
116
+ if fold >= len(candidates):
117
+ raise ValueError(
118
+ f"fold {fold} unavailable for {species_name}: "
119
+ f"only {len(candidates)} domains"
120
+ )
121
+ held_domains[species] = candidates[fold]
122
+ validation_ids.extend(groups[species, candidates[fold]])
123
+ validation_set = set(validation_ids)
124
+ training_ids = [file_id for file_id in ids if file_id not in validation_set]
125
+ training_species = {parse_file_id(file_id)[0] for file_id in training_ids}
126
+ validation_species = {parse_file_id(file_id)[0] for file_id in validation_ids}
127
+ expected_species = set(range(len(SPECIES)))
128
+ if training_species != expected_species or validation_species != expected_species:
129
+ raise ValueError("held-domain split does not cover every species")
130
+ return training_ids, validation_ids, held_domains
131
+
132
+
133
+ def load_ids(split="Training"):
134
+ """Download and read one official ID list."""
135
+ path = hf_hub_download(DATASET, f"data/metadata/{split}_ids.txt", repo_type="dataset")
136
+ with open(path, encoding="utf-8") as handle:
137
+ return [line.strip() for line in handle if line.strip()]
138
+
139
+
140
+ def load_unseen_domains():
141
+ """Return the official unseen domain index for each species index."""
142
+ path = hf_hub_download(DATASET, "data/metadata/split_summary.json", repo_type="dataset")
143
+ mapping = json.loads(Path(path).read_text(encoding="utf-8"))["unseen_domain_by_species"]
144
+ return {SPECIES.index(species): DOMAINS.index(domain) for species, domain in mapping.items()}
145
+
146
+
147
+ def download(url, path):
148
+ """Download a large file atomically with coarse progress output."""
149
+ partial = path.with_suffix(path.suffix + ".part")
150
+ with urllib.request.urlopen(url) as response, open(partial, "wb") as output:
151
+ total = int(response.headers.get("Content-Length", 0))
152
+ downloaded = 0
153
+ report_at = 0
154
+ while chunk := response.read(1 << 20):
155
+ output.write(chunk)
156
+ downloaded += len(chunk)
157
+ if downloaded >= report_at:
158
+ print(
159
+ f"downloaded {downloaded / 1e9:.1f}/{total / 1e9:.1f} GB",
160
+ flush=True,
161
+ )
162
+ report_at += 250_000_000
163
+ partial.rename(path)
164
+
165
+
166
+ def prepare_data(data):
167
+ """Download and flatten the official Zenodo development archive."""
168
+ destination = Path(data)
169
+ marker = destination / ".complete"
170
+ if marker.exists():
171
+ print(f"audio already prepared at {destination}")
172
+ return
173
+ destination.mkdir(parents=True, exist_ok=True)
174
+ archive = destination.parent / "Development_data.zip"
175
+ if not archive.exists():
176
+ download(ZENODO, archive)
177
+ with zipfile.ZipFile(archive) as source:
178
+ members = [
179
+ name
180
+ for name in source.namelist()
181
+ if name.lower().endswith(".wav")
182
+ and not name.startswith("__MACOSX")
183
+ and not Path(name).name.startswith("._")
184
+ ]
185
+ print(f"extracting {len(members)} waveforms", flush=True)
186
+ for index, name in enumerate(members, 1):
187
+ with source.open(name) as src, open(
188
+ destination / Path(name).name, "wb"
189
+ ) as dst:
190
+ shutil.copyfileobj(src, dst)
191
+ if index % 25_000 == 0:
192
+ print(f"extracted {index}/{len(members)}", flush=True)
193
+ marker.write_text(str(len(members)), encoding="utf-8")
194
+ archive.unlink()
195
+ print(f"prepared {len(members)} waveforms at {destination}")
196
+
197
+
198
+ def select_ids(ids, size, seed):
199
+ """Select a seeded species-domain-stratified subset."""
200
+ if not size or size >= len(ids):
201
+ return ids
202
+ groups = {}
203
+ for file_id in ids:
204
+ groups.setdefault(parse_file_id(file_id), []).append(file_id)
205
+ rng = random.Random(seed)
206
+ for rows in groups.values():
207
+ rng.shuffle(rows)
208
+ selected = []
209
+ while len(selected) < size:
210
+ added = False
211
+ for key in sorted(groups):
212
+ if groups[key] and len(selected) < size:
213
+ selected.append(groups[key].pop())
214
+ added = True
215
+ if not added:
216
+ break
217
+ rng.shuffle(selected)
218
+ return selected
219
+
220
+
221
+ def cpus():
222
+ """Return the CPU quota visible to this process."""
223
+ try:
224
+ quota, period = Path("/sys/fs/cgroup/cpu.max").read_text(
225
+ encoding="utf-8"
226
+ ).split()
227
+ if quota != "max":
228
+ return max(1, int(quota) // int(period))
229
+ except OSError:
230
+ pass
231
+ return len(os.sched_getaffinity(0))
232
+
233
+
234
+ def load_waveform(path, seconds=None, training=False):
235
+ """Read a mono waveform at the model sample rate, or only a random or central crop."""
236
+ start, frames = 0, -1
237
+ if seconds is not None:
238
+ info = soundfile.info(path)
239
+ frames = math.ceil(seconds * info.samplerate)
240
+ gap = max(0, info.frames - frames)
241
+ start = random.randrange(gap + 1) if training else gap // 2
242
+ waveform, sample_rate = soundfile.read(
243
+ path, start=start, frames=frames, dtype="float32", always_2d=False
244
+ )
245
+ if waveform.ndim > 1:
246
+ waveform = waveform.mean(axis=1)
247
+ if sample_rate != SAMPLE_RATE:
248
+ waveform = resample_poly(waveform, SAMPLE_RATE, sample_rate).astype(np.float32)
249
+ return waveform
250
+
251
+
252
+ def normalize_level(waveforms, level_rms):
253
+ """Scale each waveform to one RMS level so recording gain cannot reveal its domain."""
254
+ if level_rms <= 0:
255
+ return waveforms.astype(np.float32)
256
+ rms = np.sqrt(np.mean(np.square(waveforms), axis=-1, keepdims=True))
257
+ return (waveforms * (level_rms / np.maximum(rms, 1e-6))).astype(np.float32)
258
+
259
+
260
+ def split_windows(waveform, samples):
261
+ """Cover a complete clip with fixed-length windows, the last aligned to its end."""
262
+ starts = list(range(0, max(len(waveform) - samples + 1, 1), samples))
263
+ if starts[-1] + samples < len(waveform):
264
+ starts.append(len(waveform) - samples)
265
+ return np.stack(
266
+ [fixed_length(waveform[start : start + samples], samples, False) for start in starts]
267
+ )
268
+
269
+
270
+ def fixed_length(waveform, samples, training):
271
+ """Randomly or centrally crop and zero-pad a waveform."""
272
+ if len(waveform) < samples:
273
+ gap = samples - len(waveform)
274
+ left = random.randrange(gap + 1) if training else gap // 2
275
+ waveform = np.pad(waveform, (left, gap - left))
276
+ if len(waveform) > samples:
277
+ gap = len(waveform) - samples
278
+ start = random.randrange(gap + 1) if training else gap // 2
279
+ waveform = waveform[start : start + samples]
280
+ return waveform.astype(np.float32)
281
+
282
+
283
+ class MosquitoDataset(Dataset):
284
+ """Fold-filtered 8 kHz audio with species-domain sampling weights."""
285
+
286
+ def __init__(
287
+ self, root, ids, seconds, training, balance_power, synthetic=False, level_rms=LEVEL_RMS
288
+ ):
289
+ self.root = Path(root) if root else None
290
+ self.ids = ids
291
+ self.samples = round(seconds * SAMPLE_RATE)
292
+ self.training = training
293
+ self.synthetic = synthetic
294
+ self.level_rms = level_rms
295
+ self.labels = [parse_file_id(file_id) for file_id in ids]
296
+ counts = Counter(self.labels)
297
+ self.sample_weights = [counts[label] ** -balance_power for label in self.labels]
298
+
299
+ def __len__(self):
300
+ return len(self.ids)
301
+
302
+ def _synthetic_waveform(self, index):
303
+ species, domain = self.labels[index]
304
+ row = int(self.ids[index].rsplit("_", maxsplit=1)[1])
305
+ time = np.arange(self.samples, dtype=np.float32) / SAMPLE_RATE
306
+ phase = (row % 7) * math.pi / 7
307
+ return 0.2 * np.sin(2 * np.pi * (180 + 35 * species + domain) * time + phase)
308
+
309
+ def __getitem__(self, index):
310
+ if self.synthetic:
311
+ waveform = self._synthetic_waveform(index)
312
+ else:
313
+ waveform = load_waveform(
314
+ self.root / f"{self.ids[index]}.wav",
315
+ self.samples / SAMPLE_RATE,
316
+ self.training,
317
+ )
318
+ waveform = fixed_length(waveform, self.samples, self.training)
319
+ waveform = normalize_level(waveform, self.level_rms)
320
+ if self.training:
321
+ waveform = np.roll(waveform, random.randrange(len(waveform)))
322
+ waveform = waveform * 10 ** random.uniform(-0.3, 0.3)
323
+ waveform = waveform + np.random.normal(
324
+ 0, random.uniform(0, 0.005), len(waveform)
325
+ )
326
+ species, domain = self.labels[index]
327
+ return {
328
+ "file_id": self.ids[index],
329
+ "waveform": waveform.astype(np.float32),
330
+ "species": species,
331
+ "domain": domain,
332
+ }
333
+
334
+
335
+ class AudioCollator:
336
+ """Stack fixed-length waveforms and labels."""
337
+
338
+ def __call__(self, rows):
339
+ waveforms = torch.from_numpy(np.stack([row["waveform"] for row in rows]))
340
+ return {
341
+ "file_ids": [row["file_id"] for row in rows],
342
+ "waveforms": waveforms,
343
+ "lengths": torch.full((len(rows),), waveforms.size(1), dtype=torch.long),
344
+ "species": torch.tensor([row["species"] for row in rows]),
345
+ "domains": torch.tensor([row["domain"] for row in rows]),
346
+ }
347
+
348
+
349
+ class ClipDataset(MosquitoDataset):
350
+ """Complete clips split into training-length evaluation windows."""
351
+
352
+ def __getitem__(self, index):
353
+ if self.synthetic:
354
+ waveform = self._synthetic_waveform(index)
355
+ else:
356
+ waveform = load_waveform(self.root / f"{self.ids[index]}.wav")
357
+ species, domain = self.labels[index]
358
+ return {
359
+ "windows": normalize_level(split_windows(waveform, self.samples), self.level_rms),
360
+ "species": species,
361
+ "domain": domain,
362
+ }
363
+
364
+
365
+ class WindowCollator:
366
+ """Concatenate every clip's windows and record which clip owns each."""
367
+
368
+ def __call__(self, rows):
369
+ counts = torch.tensor([len(row["windows"]) for row in rows])
370
+ waveforms = torch.from_numpy(np.concatenate([row["windows"] for row in rows]))
371
+ return {
372
+ "waveforms": waveforms,
373
+ "lengths": torch.full((len(waveforms),), waveforms.size(1), dtype=torch.long),
374
+ "owners": torch.repeat_interleave(torch.arange(len(rows)), counts),
375
+ "species": torch.tensor([row["species"] for row in rows]),
376
+ "domains": torch.tensor([row["domain"] for row in rows]),
377
+ }
378
+
379
+
380
+ def synthetic_ids(rows_per_cell):
381
+ """Build official-shaped IDs with the real species-domain availability pattern."""
382
+ return [
383
+ f"S_{species + 1}_D_{domain + 1}_{row + 1}"
384
+ for species, domains in enumerate(SYNTHETIC_DOMAINS)
385
+ for domain in domains
386
+ for row in range(rows_per_cell)
387
+ ]
388
+
389
+
390
+ def make_datasets(args):
391
+ """Create one pseudo-unseen fold, or for fold -1 all training rows and the official
392
+ validation split."""
393
+ ids = synthetic_ids(args.synthetic_rows) if args.synthetic else load_ids()
394
+ if args.fold < 0:
395
+ training_ids = ids
396
+ validation_ids = ids if args.synthetic else load_ids("Validation")
397
+ held_domains = dict.fromkeys(range(len(SPECIES)))
398
+ else:
399
+ training_ids, validation_ids, held_domains = held_domain_split(
400
+ ids, args.fold, seed=args.seed
401
+ )
402
+ training_ids = select_ids(training_ids, args.train_size, args.seed)
403
+ validation_ids = select_ids(validation_ids, args.validation_size, args.seed)
404
+ training = MosquitoDataset(
405
+ args.data,
406
+ training_ids,
407
+ args.seconds,
408
+ True,
409
+ args.balance_power,
410
+ args.synthetic,
411
+ args.level_rms,
412
+ )
413
+ validation = MosquitoDataset(
414
+ args.data,
415
+ validation_ids,
416
+ args.seconds,
417
+ False,
418
+ args.balance_power,
419
+ args.synthetic,
420
+ args.level_rms,
421
+ )
422
+ expected = set(range(len(SPECIES)))
423
+ if {label[0] for label in training.labels} != expected:
424
+ raise ValueError("training size limit removed a species")
425
+ if {label[0] for label in validation.labels} != expected:
426
+ raise ValueError("validation size limit removed a species")
427
+ return training, validation, held_domains
428
+
429
+
430
+ def seed_worker(_worker_id):
431
+ """Seed NumPy and Python from a DataLoader worker's PyTorch seed."""
432
+ worker_seed = torch.initial_seed() % (2**32)
433
+ random.seed(worker_seed)
434
+ np.random.seed(worker_seed)
435
+
436
+
437
+ def make_loader(dataset, batch_size, workers, training, seed, collator=None):
438
+ """Create a deterministic loader, balanced by species-domain cell for training."""
439
+ generator = torch.Generator().manual_seed(seed)
440
+ sampler = None
441
+ if training:
442
+ sampler = WeightedRandomSampler(
443
+ dataset.sample_weights,
444
+ num_samples=len(dataset),
445
+ replacement=True,
446
+ generator=generator,
447
+ )
448
+ return DataLoader(
449
+ dataset,
450
+ batch_size=batch_size,
451
+ sampler=sampler,
452
+ shuffle=False,
453
+ num_workers=workers,
454
+ pin_memory=torch.cuda.is_available(),
455
+ persistent_workers=workers > 0,
456
+ worker_init_fn=seed_worker,
457
+ generator=generator,
458
+ collate_fn=collator or AudioCollator(),
459
+ )
460
+
461
+
462
+ def deterministic_view(dataset, ids=None):
463
+ """Return an unaugmented view over a dataset's IDs."""
464
+ return MosquitoDataset(
465
+ dataset.root,
466
+ dataset.ids if ids is None else ids,
467
+ dataset.samples / SAMPLE_RATE,
468
+ training=False,
469
+ balance_power=0.0,
470
+ synthetic=dataset.synthetic,
471
+ level_rms=dataset.level_rms,
472
+ )
473
+
474
+
475
+ @torch.no_grad()
476
+ def compute_normalization(dataset, size, batch_size, workers, device, seed, min_hz):
477
+ """Estimate per-mel mean and variance from fold-training rows only."""
478
+ ids = select_ids(dataset.ids, size, seed)
479
+ loader = make_loader(
480
+ deterministic_view(dataset, ids), batch_size, workers, False, seed
481
+ )
482
+ frontend = LogMelFrontend(min_hz=min_hz).to(device)
483
+ total = torch.zeros(N_MELS, dtype=torch.float64, device=device)
484
+ squared = torch.zeros_like(total)
485
+ frames = 0
486
+ for batch in loader:
487
+ waveforms = batch["waveforms"].to(device, non_blocking=True)
488
+ lengths = batch["lengths"].to(device, non_blocking=True)
489
+ features, frame_lengths = frontend(waveforms, lengths)
490
+ positions = torch.arange(features.size(1), device=device)[None, :]
491
+ mask = positions < frame_lengths[:, None]
492
+ selected = features[mask].double()
493
+ total += selected.sum(0)
494
+ squared += selected.square().sum(0)
495
+ frames += selected.size(0)
496
+ mean = total / frames
497
+ variance = squared / frames - mean.square()
498
+ return mean.float().cpu(), variance.clamp_min(1e-6).sqrt().float().cpu()
499
+
500
+
501
+ def balanced_accuracy(gold, predicted):
502
+ """Return mean per-species recall over labels present in gold."""
503
+ gold = np.asarray(gold)
504
+ predicted = np.asarray(predicted)
505
+ recalls = [
506
+ np.mean(predicted[gold == label] == label)
507
+ for label in sorted(set(gold.tolist()))
508
+ ]
509
+ return float(np.mean(recalls)) if recalls else 0.0
510
+
511
+
512
+ def prediction_metrics(labels, predictions):
513
+ """Return balanced accuracy, per-species recall and a gold-by-predicted confusion."""
514
+ labels = labels.cpu().numpy()
515
+ predictions = predictions.cpu().numpy()
516
+ confusion = np.bincount(
517
+ labels * len(SPECIES) + predictions, minlength=len(SPECIES) ** 2
518
+ ).reshape(len(SPECIES), len(SPECIES))
519
+ return {
520
+ "balanced_accuracy": balanced_accuracy(labels, predictions),
521
+ "rows": len(labels),
522
+ "per_species": {
523
+ SPECIES[label]: float(np.mean(predictions[labels == label] == label))
524
+ for label in sorted(set(labels.tolist()))
525
+ },
526
+ "confusion": confusion.tolist(),
527
+ }
528
+
529
+
530
+ def official_metrics(labels, predictions, domains, unseen_domains):
531
+ """Score the BioDCASE seen and unseen species-domain partitions."""
532
+ unseen = torch.tensor(
533
+ [
534
+ unseen_domains[label.item()] == domain.item()
535
+ for label, domain in zip(labels, domains)
536
+ ],
537
+ dtype=torch.bool,
538
+ )
539
+ return {
540
+ "seen": prediction_metrics(labels[~unseen], predictions[~unseen]),
541
+ "unseen": prediction_metrics(labels[unseen], predictions[unseen]),
542
+ }
543
+
544
+
545
+ def fold_prediction_metrics(labels, predictions, domains, held_domains):
546
+ """Split fold metrics into true domain holdouts and singleton reserves."""
547
+ metrics = prediction_metrics(labels, predictions)
548
+ domain_holdout = torch.tensor(
549
+ [held_domains[label.item()] is not None for label in labels],
550
+ dtype=torch.bool,
551
+ )
552
+ metrics["domain_holdout"] = prediction_metrics(
553
+ labels[domain_holdout], predictions[domain_holdout]
554
+ )
555
+ metrics["same_domain"] = prediction_metrics(
556
+ labels[~domain_holdout], predictions[~domain_holdout]
557
+ )
558
+ expected_domains = torch.tensor(
559
+ [
560
+ held_domains[label.item()]
561
+ if held_domains[label.item()] is not None
562
+ else domain.item()
563
+ for label, domain in zip(labels, domains)
564
+ ]
565
+ )
566
+ if not torch.equal(domains, expected_domains):
567
+ raise ValueError("validation rows do not match the declared fold policy")
568
+ return metrics
569
+
570
+
571
+ @torch.no_grad()
572
+ def collect_outputs(model, frontend, dataset, batch_size, workers, device, seed):
573
+ """Collect deterministic logits and embeddings for one dataset."""
574
+ model.eval()
575
+ loader = make_loader(dataset, batch_size, workers, False, seed)
576
+ collected = {"labels": [], "domains": [], "logits": [], "embeddings": []}
577
+ for batch in loader:
578
+ waveforms = batch["waveforms"].to(device, non_blocking=True)
579
+ lengths = batch["lengths"].to(device, non_blocking=True)
580
+ domains = batch["domains"].to(device, non_blocking=True)
581
+ features, frame_lengths = frontend(waveforms, lengths)
582
+ outputs = model(features, frame_lengths, domains)
583
+ collected["labels"].append(batch["species"])
584
+ collected["domains"].append(batch["domains"])
585
+ collected["logits"].append(outputs["species_logits"].cpu())
586
+ collected["embeddings"].append(outputs["embedding"].cpu())
587
+ return {name: torch.cat(values) for name, values in collected.items()}
588
+
589
+
590
+ @torch.no_grad()
591
+ def collect_clip_outputs(model, frontend, dataset, batch_size, workers, device, seed):
592
+ """Average window probabilities and embeddings over each complete clip."""
593
+ model.eval()
594
+ loader = make_loader(dataset, batch_size, workers, False, seed, WindowCollator())
595
+ collected = {"labels": [], "domains": [], "probabilities": [], "embeddings": []}
596
+ for batch in loader:
597
+ owners = batch["owners"].to(device, non_blocking=True)
598
+ domains = batch["domains"].to(device, non_blocking=True)
599
+ window_outputs = {"probabilities": [], "embeddings": []}
600
+ for start in range(0, len(owners), batch_size):
601
+ window = slice(start, start + batch_size)
602
+ features, frame_lengths = frontend(
603
+ batch["waveforms"][window].to(device, non_blocking=True),
604
+ batch["lengths"][window].to(device, non_blocking=True),
605
+ )
606
+ outputs = model(features, frame_lengths, domains[owners[window]])
607
+ window_outputs["probabilities"].append(
608
+ functional.softmax(outputs["species_logits"], dim=1)
609
+ )
610
+ window_outputs["embeddings"].append(outputs["embedding"])
611
+ counts = torch.bincount(owners, minlength=len(domains)).unsqueeze(1)
612
+ for name, values in window_outputs.items():
613
+ values = torch.cat(values)
614
+ totals = values.new_zeros(len(domains), values.size(1))
615
+ collected[name].append((totals.index_add_(0, owners, values) / counts).cpu())
616
+ collected["labels"].append(batch["species"])
617
+ collected["domains"].append(batch["domains"])
618
+ return {name: torch.cat(values) for name, values in collected.items()}
619
+
620
+
621
+ def mahalanobis_fit(reference_embeddings, reference_labels):
622
+ """Return class means and the precision of one regularized shared covariance matrix."""
623
+ if set(reference_labels.tolist()) != set(range(len(SPECIES))):
624
+ raise ValueError("Mahalanobis reference data must cover every species")
625
+ means = torch.stack(
626
+ [
627
+ reference_embeddings[reference_labels == label].mean(0)
628
+ for label in range(len(SPECIES))
629
+ ]
630
+ )
631
+ residuals = reference_embeddings - means[reference_labels]
632
+ covariance = residuals.T @ residuals / max(1, len(residuals) - len(SPECIES))
633
+ regularization = covariance.diagonal().mean().clamp_min(1e-6) * 1e-3
634
+ covariance = covariance + regularization * torch.eye(covariance.size(0))
635
+ return means, torch.linalg.pinv(covariance)
636
+
637
+
638
+ def mahalanobis_predict(means, precision, embeddings):
639
+ """Assign each embedding to the nearest class mean in Mahalanobis distance."""
640
+ differences = embeddings[:, None, :] - means[None, :, :]
641
+ distances = torch.einsum("ncd,de,nce->nc", differences, precision, differences)
642
+ return distances.argmin(1)
643
+
644
+
645
+ def sampled_priors(dataset, device):
646
+ """Return species priors under the balanced training sampler for logit adjustment."""
647
+ totals = torch.zeros(len(SPECIES), dtype=torch.float64).index_add_(
648
+ 0,
649
+ torch.tensor([species for species, _domain in dataset.labels]),
650
+ torch.tensor(dataset.sample_weights, dtype=torch.float64),
651
+ )
652
+ return (totals / totals.sum()).float().to(device)
653
+
654
+
655
+ def grl_strength(step, total_steps):
656
+ """Return the standard sigmoid gradient-reversal warmup."""
657
+ progress = step / max(1, total_steps - 1)
658
+ return 2 / (1 + math.exp(-10 * progress)) - 1
659
+
660
+
661
+ def learning_rate_schedule(total_steps, warmup_ratio):
662
+ """Return a linear-warmup cosine-decay multiplier."""
663
+ warmup_steps = max(1, round(total_steps * warmup_ratio))
664
+
665
+ def multiplier(step):
666
+ if step < warmup_steps:
667
+ return (step + 1) / warmup_steps
668
+ progress = (step - warmup_steps) / max(1, total_steps - warmup_steps)
669
+ return 0.5 * (1 + math.cos(math.pi * min(1.0, progress)))
670
+
671
+ return multiplier
672
+
673
+
674
+ def train_epoch(model, frontend, loader, optimizer, scheduler, priors, epoch, state):
675
+ """Run one balanced training epoch or stop at the configured step limit."""
676
+ model.train()
677
+ totals = {}
678
+ batches = 0
679
+ for batch in loader:
680
+ if 0 < state["max_steps"] <= state["step"]:
681
+ break
682
+ waveforms = batch["waveforms"].to(state["device"], non_blocking=True)
683
+ lengths = batch["lengths"].to(state["device"], non_blocking=True)
684
+ species = batch["species"].to(state["device"], non_blocking=True)
685
+ domains = batch["domains"].to(state["device"], non_blocking=True)
686
+ with torch.no_grad():
687
+ features, frame_lengths = frontend(waveforms, lengths)
688
+ optimizer.zero_grad(set_to_none=True)
689
+ outputs = model(
690
+ features,
691
+ frame_lengths,
692
+ domains,
693
+ domain_strength=grl_strength(state["step"], state["total_steps"]),
694
+ augment=True,
695
+ )
696
+ losses = training_loss(outputs, species, domains, priors, epoch + 1)
697
+ losses["loss"].backward()
698
+ nn.utils.clip_grad_norm_(model.parameters(), state["gradient_clip"])
699
+ optimizer.step()
700
+ scheduler.step()
701
+ state["step"] += 1
702
+ batches += 1
703
+ for name, value in losses.items():
704
+ totals[name] = totals.get(name, 0.0) + value.detach().item()
705
+ if not batches:
706
+ return None
707
+ return {name: value / batches for name, value in totals.items()}
708
+
709
+
710
+ def checkpoint_state(model, frontend, optimizer, scheduler, epoch, state, metrics, args):
711
+ """Build a resumable best-checkpoint payload."""
712
+ return {
713
+ "model": model.state_dict(),
714
+ "frontend": frontend.state_dict(),
715
+ "optimizer": optimizer.state_dict(),
716
+ "scheduler": scheduler.state_dict(),
717
+ "epoch": epoch,
718
+ "step": state["step"],
719
+ "metrics": metrics,
720
+ "arguments": {key: value for key, value in vars(args).items() if key != "run"},
721
+ }
722
+
723
+
724
+ def final_metrics(model, frontend, training, validation, held_domains, args, device):
725
+ """Score softmax and shared-covariance Mahalanobis inference."""
726
+ reference_ids = select_ids(training.ids, args.reference_size, args.seed + 1)
727
+ reference = deterministic_view(training, reference_ids)
728
+ reference_outputs = collect_outputs(
729
+ model,
730
+ frontend,
731
+ reference,
732
+ args.eval_batch_size,
733
+ args.workers,
734
+ device,
735
+ args.seed,
736
+ )
737
+ validation_outputs = collect_outputs(
738
+ model,
739
+ frontend,
740
+ validation,
741
+ args.eval_batch_size,
742
+ args.workers,
743
+ device,
744
+ args.seed,
745
+ )
746
+ softmax = fold_prediction_metrics(
747
+ validation_outputs["labels"],
748
+ validation_outputs["logits"].argmax(1),
749
+ validation_outputs["domains"],
750
+ held_domains,
751
+ )
752
+ mahalanobis = fold_prediction_metrics(
753
+ validation_outputs["labels"],
754
+ mahalanobis_predict(
755
+ *mahalanobis_fit(reference_outputs["embeddings"], reference_outputs["labels"]),
756
+ validation_outputs["embeddings"],
757
+ ),
758
+ validation_outputs["domains"],
759
+ held_domains,
760
+ )
761
+ return softmax, mahalanobis, reference_outputs
762
+
763
+
764
+ def train(args):
765
+ """Train and select MTRCNN-DG on one training-only pseudo-unseen fold."""
766
+ set_seed(args.seed)
767
+ if not args.synthetic and not (Path(args.data) / ".complete").exists():
768
+ prepare_data(args.data)
769
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
770
+ training, validation, held_domains = make_datasets(args)
771
+ mean, std = compute_normalization(
772
+ training,
773
+ args.stats_size,
774
+ args.eval_batch_size,
775
+ args.workers,
776
+ device,
777
+ args.seed,
778
+ args.min_hz,
779
+ )
780
+ frontend = LogMelFrontend(mean, std, args.min_hz).to(device)
781
+ model = MTRCNNDG(dropout=args.dropout).to(device)
782
+ loader = make_loader(
783
+ training, args.batch_size, args.workers, True, args.seed
784
+ )
785
+ natural_steps = args.epochs * len(loader)
786
+ total_steps = min(natural_steps, args.max_steps) if args.max_steps > 0 else natural_steps
787
+ optimizer = torch.optim.AdamW(
788
+ model.parameters(), lr=args.learning_rate, weight_decay=args.weight_decay
789
+ )
790
+ scheduler = torch.optim.lr_scheduler.LambdaLR(
791
+ optimizer, learning_rate_schedule(total_steps, args.warmup_ratio)
792
+ )
793
+ priors = sampled_priors(training, device)
794
+ output = Path(args.output)
795
+ output.mkdir(parents=True, exist_ok=True)
796
+ checkpoint = output / "checkpoint.pt"
797
+ state = {
798
+ "device": device,
799
+ "step": 0,
800
+ "total_steps": total_steps,
801
+ "max_steps": args.max_steps,
802
+ "gradient_clip": args.gradient_clip,
803
+ }
804
+ history = []
805
+ best_score = -1.0
806
+ best_epoch = 0
807
+ stale_epochs = 0
808
+ for epoch in range(args.epochs):
809
+ losses = train_epoch(
810
+ model, frontend, loader, optimizer, scheduler, priors, epoch, state
811
+ )
812
+ if losses is None:
813
+ break
814
+ validation_outputs = collect_outputs(
815
+ model,
816
+ frontend,
817
+ validation,
818
+ args.eval_batch_size,
819
+ args.workers,
820
+ device,
821
+ args.seed,
822
+ )
823
+ metrics = fold_prediction_metrics(
824
+ validation_outputs["labels"],
825
+ validation_outputs["logits"].argmax(1),
826
+ validation_outputs["domains"],
827
+ held_domains,
828
+ )
829
+ record = {"epoch": epoch + 1, "step": state["step"], "losses": losses, **metrics}
830
+ history.append(record)
831
+ print(json.dumps(record), flush=True)
832
+ # Without a held-out domain (fold -1) the last epoch is kept.
833
+ score = metrics["domain_holdout"]["balanced_accuracy"] if args.fold >= 0 else epoch
834
+ if score > best_score:
835
+ best_score = score
836
+ best_epoch = epoch + 1
837
+ stale_epochs = 0
838
+ torch.save(
839
+ checkpoint_state(
840
+ model,
841
+ frontend,
842
+ optimizer,
843
+ scheduler,
844
+ epoch + 1,
845
+ state,
846
+ metrics,
847
+ args,
848
+ ),
849
+ checkpoint,
850
+ )
851
+ else:
852
+ stale_epochs += 1
853
+ reached_limit = 0 < args.max_steps <= state["step"]
854
+ if reached_limit or stale_epochs >= args.patience:
855
+ break
856
+ saved = torch.load(checkpoint, map_location=device, weights_only=True)
857
+ model.load_state_dict(saved["model"])
858
+ frontend.load_state_dict(saved["frontend"])
859
+ softmax, mahalanobis, reference = final_metrics(
860
+ model, frontend, training, validation, held_domains, args, device
861
+ )
862
+ result = {
863
+ "device": str(device),
864
+ "fold": args.fold,
865
+ "seed": args.seed,
866
+ "training_rows": len(training),
867
+ "validation_rows": len(validation),
868
+ "normalization_rows": len(select_ids(training.ids, args.stats_size, args.seed)),
869
+ "reference_rows": len(reference["labels"]),
870
+ "held_domains": {
871
+ SPECIES[species]: DOMAINS[domain] if domain is not None else "same-domain"
872
+ for species, domain in held_domains.items()
873
+ },
874
+ "best_epoch": best_epoch,
875
+ "steps": state["step"],
876
+ "softmax": softmax,
877
+ "mahalanobis": mahalanobis,
878
+ "history": history,
879
+ "checkpoint": str(checkpoint),
880
+ "test": test_metrics(model, frontend, reference, args, device) if args.test else None,
881
+ }
882
+ (output / "results.json").write_text(
883
+ json.dumps(result, indent=1), encoding="utf-8"
884
+ )
885
+ print(json.dumps(result, indent=1))
886
+
887
+
888
+ def test_metrics(model, frontend, reference, args, device):
889
+ """Score the official development test split with softmax and Mahalanobis inference."""
890
+ test_ids = (
891
+ [
892
+ f"S_{species + 1}_D_{domain + 1}_{row + 1}"
893
+ for species in range(len(SPECIES))
894
+ for domain in range(len(DOMAINS))
895
+ for row in range(args.synthetic_rows)
896
+ ]
897
+ if args.synthetic
898
+ else load_ids("Test")
899
+ )
900
+ test = ClipDataset(
901
+ args.data,
902
+ select_ids(test_ids, args.test_size, args.seed),
903
+ args.seconds,
904
+ False,
905
+ 0.0,
906
+ args.synthetic,
907
+ args.level_rms,
908
+ )
909
+ outputs = collect_clip_outputs(
910
+ model, frontend, test, args.eval_batch_size, args.workers, device, args.seed
911
+ )
912
+ unseen_domains = load_unseen_domains()
913
+ mahalanobis = mahalanobis_predict(
914
+ *mahalanobis_fit(reference["embeddings"], reference["labels"]), outputs["embeddings"]
915
+ )
916
+ return {
917
+ "rows": len(test),
918
+ "softmax": official_metrics(
919
+ outputs["labels"],
920
+ outputs["probabilities"].argmax(1),
921
+ outputs["domains"],
922
+ unseen_domains,
923
+ ),
924
+ "mahalanobis": official_metrics(
925
+ outputs["labels"], mahalanobis, outputs["domains"], unseen_domains
926
+ ),
927
+ }
928
+
929
+
930
+ def evaluate(args):
931
+ """Score a checkpoint on its validation rows and on the official development test split."""
932
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
933
+ saved = torch.load(args.checkpoint, map_location=device, weights_only=True)
934
+ options = argparse.Namespace(**saved["arguments"])
935
+ options.data = args.data
936
+ options.eval_batch_size = args.eval_batch_size
937
+ options.workers = args.workers
938
+ options.test_size = args.test_size
939
+ set_seed(options.seed)
940
+ if not options.synthetic and not (Path(args.data) / ".complete").exists():
941
+ prepare_data(args.data)
942
+ training, validation, held_domains = make_datasets(options)
943
+ frontend = LogMelFrontend(min_hz=options.min_hz).to(device)
944
+ frontend.load_state_dict(saved["frontend"])
945
+ model = MTRCNNDG(dropout=options.dropout).to(device)
946
+ model.load_state_dict(saved["model"])
947
+ fold_softmax, fold_mahalanobis, reference = final_metrics(
948
+ model, frontend, training, validation, held_domains, options, device
949
+ )
950
+ result = {
951
+ "checkpoint": args.checkpoint,
952
+ "epoch": saved["epoch"],
953
+ "fold": options.fold,
954
+ "seed": options.seed,
955
+ "arguments": saved["arguments"],
956
+ "test": test_metrics(model, frontend, reference, options, device),
957
+ "fold_softmax": fold_softmax,
958
+ "fold_mahalanobis": fold_mahalanobis,
959
+ }
960
+ output = Path(args.output)
961
+ output.parent.mkdir(parents=True, exist_ok=True)
962
+ output.write_text(json.dumps(result, indent=1), encoding="utf-8")
963
+ print(json.dumps(result, indent=1))
964
+ if args.export:
965
+ export(model, frontend, reference, saved["arguments"], result, Path(args.export))
966
+ if args.repo:
967
+ HfApi().upload_folder(
968
+ repo_id=args.repo,
969
+ folder_path=args.export,
970
+ commit_message="Add MTRCNN-DG weights and evaluation",
971
+ )
972
+
973
+
974
+ def export(model, frontend, reference, arguments, result, folder):
975
+ """Write weights, Mahalanobis statistics, this script and the evaluation to folder."""
976
+ means, precision = mahalanobis_fit(reference["embeddings"], reference["labels"])
977
+ (folder / "results").mkdir(parents=True, exist_ok=True)
978
+ torch.save(
979
+ {
980
+ "model": model.state_dict(),
981
+ "frontend": frontend.state_dict(),
982
+ "arguments": arguments,
983
+ "means": means,
984
+ "precision": precision,
985
+ },
986
+ folder / "model.pt",
987
+ )
988
+ shutil.copy(__file__, folder / "mtrcnn.py")
989
+ (folder / "results" / "eval.json").write_text(json.dumps(result, indent=1), encoding="utf-8")
990
+
991
+
992
+ def predict(args):
993
+ """Print the Mahalanobis species prediction for each WAV file."""
994
+ path = Path(args.model)
995
+ if not path.is_file():
996
+ path = hf_hub_download(args.model, "model.pt")
997
+ saved = torch.load(path, map_location="cpu", weights_only=True)
998
+ options = saved["arguments"]
999
+ frontend = LogMelFrontend(min_hz=options["min_hz"])
1000
+ frontend.load_state_dict(saved["frontend"])
1001
+ model = MTRCNNDG().eval()
1002
+ model.load_state_dict(saved["model"])
1003
+ samples = round(options["seconds"] * SAMPLE_RATE)
1004
+ for file in args.files:
1005
+ windows = split_windows(load_waveform(file), samples)
1006
+ waveforms = torch.from_numpy(normalize_level(windows, options["level_rms"]))
1007
+ with torch.no_grad():
1008
+ features, lengths = frontend(waveforms, torch.full((len(waveforms),), samples))
1009
+ embedding = model(features, lengths, torch.zeros(len(waveforms), dtype=torch.long))[
1010
+ "embedding"
1011
+ ].mean(0, keepdim=True)
1012
+ species = mahalanobis_predict(saved["means"], saved["precision"], embedding).item()
1013
+ print(json.dumps({"file": file, "species": SPECIES[species]}))
1014
+
1015
+
1016
+ def card(args):
1017
+ """Render card.jinja into card/README.md with the results stored in the model repo."""
1018
+ try:
1019
+ folder = Path(snapshot_download(args.repo, allow_patterns="results/*.json"))
1020
+ paths = folder.glob("results/*.json")
1021
+ except RepositoryNotFoundError:
1022
+ paths = []
1023
+ results = {path.stem: json.loads(path.read_text(encoding="utf-8")) for path in paths}
1024
+ runs = {}
1025
+ for name, result in sorted(results.items()):
1026
+ if name.startswith("train-"):
1027
+ runs.setdefault(name.split("-")[1], []).append(result["test"])
1028
+ seeds = {
1029
+ config: {
1030
+ inference: {
1031
+ partition: [run[inference][partition]["balanced_accuracy"] * 100 for run in tests]
1032
+ for partition in ("seen", "unseen")
1033
+ }
1034
+ for inference in ("softmax", "mahalanobis")
1035
+ }
1036
+ for config, tests in runs.items()
1037
+ }
1038
+ evaluation = results.get("eval")
1039
+ data = ModelCardData(
1040
+ model_name=args.repo.split("/")[1],
1041
+ datasets=[DATASET],
1042
+ license="cc-by-4.0",
1043
+ tags=["audio-classification", "bioacoustics", "mosquito", "domain-generalization"],
1044
+ eval_results=[
1045
+ EvalResult(
1046
+ task_type="audio-classification",
1047
+ dataset_type=DATASET,
1048
+ dataset_name="BioDCASE 2026 Task 5 development test, unseen domains",
1049
+ metric_type="balanced_accuracy",
1050
+ metric_name="Unseen-domain balanced accuracy",
1051
+ metric_value=round(
1052
+ evaluation["test"]["mahalanobis"]["unseen"]["balanced_accuracy"] * 100, 2
1053
+ ),
1054
+ )
1055
+ ]
1056
+ if evaluation
1057
+ else None,
1058
+ )
1059
+ rendered = ModelCard.from_template(
1060
+ data,
1061
+ template_path=Path(__file__).parent / "card.jinja",
1062
+ repo=args.repo,
1063
+ dataset=DATASET,
1064
+ species=SPECIES,
1065
+ evaluation=evaluation,
1066
+ seeds=seeds,
1067
+ mean=statistics.mean,
1068
+ stdev=statistics.stdev,
1069
+ )
1070
+ output = Path(__file__).parent / "card"
1071
+ output.mkdir(exist_ok=True)
1072
+ rendered.save(output / "README.md")
1073
+
1074
+
1075
+ def set_seed(seed):
1076
+ """Seed Python, NumPy and PyTorch."""
1077
+ random.seed(seed)
1078
+ np.random.seed(seed)
1079
+ torch.manual_seed(seed)
1080
+
1081
+
1082
+ def mel_filter_bank(min_hz):
1083
+ """Return a triangular Slaney-style mel filter bank from min_hz to Nyquist."""
1084
+ frequencies = torch.linspace(0, SAMPLE_RATE / 2, N_FFT // 2 + 1)
1085
+ mel_min = 2595 * math.log10(1 + min_hz / 700)
1086
+ mel_max = 2595 * math.log10(1 + (SAMPLE_RATE / 2) / 700)
1087
+ mel_points = torch.linspace(mel_min, mel_max, N_MELS + 2)
1088
+ hz_points = 700 * (torch.pow(10, mel_points / 2595) - 1)
1089
+ lower = (frequencies[:, None] - hz_points[:-2]) / (
1090
+ hz_points[1:-1] - hz_points[:-2]
1091
+ )
1092
+ upper = (hz_points[2:] - frequencies[:, None]) / (
1093
+ hz_points[2:] - hz_points[1:-1]
1094
+ )
1095
+ return torch.minimum(lower, upper).clamp_min(0)
1096
+
1097
+
1098
+ class LogMelFrontend(nn.Module):
1099
+ """Fixed 8 kHz, 64-bin log-mel frontend with global normalization."""
1100
+
1101
+ def __init__(self, mean=None, std=None, min_hz=MIN_HZ):
1102
+ super().__init__()
1103
+ self.register_buffer("window", torch.hann_window(N_FFT), persistent=False)
1104
+ self.register_buffer("mel_filters", mel_filter_bank(min_hz), persistent=False)
1105
+ self.register_buffer(
1106
+ "mean",
1107
+ torch.zeros(N_MELS) if mean is None else torch.as_tensor(mean).float(),
1108
+ )
1109
+ self.register_buffer(
1110
+ "std",
1111
+ torch.ones(N_MELS) if std is None else torch.as_tensor(std).float(),
1112
+ )
1113
+
1114
+ def forward(self, waveforms, sample_lengths):
1115
+ """Convert padded waveforms to normalized log-mel frames."""
1116
+ spectrum = torch.stft(
1117
+ waveforms,
1118
+ n_fft=N_FFT,
1119
+ hop_length=HOP_LENGTH,
1120
+ win_length=N_FFT,
1121
+ window=self.window,
1122
+ center=True,
1123
+ pad_mode="reflect",
1124
+ return_complex=True,
1125
+ )
1126
+ power = spectrum.abs().square().transpose(1, 2)
1127
+ features = 10 * torch.log10((power @ self.mel_filters).clamp_min(1e-10))
1128
+ features = (features - self.mean) / self.std.clamp_min(1e-6)
1129
+ frame_lengths = torch.div(sample_lengths, HOP_LENGTH, rounding_mode="floor") + 1
1130
+ return features, frame_lengths
1131
+
1132
+
1133
+ def different_domain_partners(domains):
1134
+ """Choose a random different-domain partner for every available sample."""
1135
+ partners = []
1136
+ for index, domain in enumerate(domains):
1137
+ candidates = torch.nonzero(domains != domain, as_tuple=False).flatten()
1138
+ if candidates.numel():
1139
+ choice = torch.randint(candidates.numel(), (), device=domains.device)
1140
+ partners.append(candidates[choice])
1141
+ else:
1142
+ partners.append(torch.tensor(index, device=domains.device))
1143
+ return torch.stack(partners)
1144
+
1145
+
1146
+ class FourierMix(nn.Module):
1147
+ """Mix log-mel Fourier amplitudes between recording domains."""
1148
+
1149
+ def __init__(self, probability=0.5):
1150
+ super().__init__()
1151
+ self.probability = probability
1152
+
1153
+ def forward(self, features, domains):
1154
+ """Optionally replace each sample's amplitude style across domains."""
1155
+ if not self.training or torch.rand((), device=features.device) >= self.probability:
1156
+ return features
1157
+ partners = different_domain_partners(domains)
1158
+ spectrum = torch.fft.rfft2(features, dim=(-2, -1), norm="ortho")
1159
+ amplitude = spectrum.abs()
1160
+ phase = spectrum / amplitude.clamp_min(1e-8)
1161
+ weight = torch.rand(features.size(0), 1, 1, device=features.device)
1162
+ mixed_amplitude = (1 - weight) * amplitude + weight * amplitude[partners]
1163
+ return torch.fft.irfft2(
1164
+ mixed_amplitude * phase,
1165
+ s=features.shape[-2:],
1166
+ dim=(-2, -1),
1167
+ norm="ortho",
1168
+ )
1169
+
1170
+
1171
+ class MixStyle(nn.Module):
1172
+ """Mix channel statistics between domains after a convolutional stage."""
1173
+
1174
+ def __init__(self, probability=0.5, alpha=0.1):
1175
+ super().__init__()
1176
+ self.probability = probability
1177
+ self.beta = torch.distributions.Beta(alpha, alpha)
1178
+
1179
+ def forward(self, features, domains):
1180
+ """Optionally mix feature statistics across recording domains."""
1181
+ if not self.training or torch.rand((), device=features.device) >= self.probability:
1182
+ return features
1183
+ mean = features.mean(dim=(2, 3), keepdim=True).detach()
1184
+ std = (features.var(dim=(2, 3), keepdim=True, unbiased=False) + 1e-6).sqrt().detach()
1185
+ normalized = (features - mean) / std
1186
+ partners = different_domain_partners(domains)
1187
+ weight = self.beta.sample((features.size(0), 1, 1, 1)).to(features.device)
1188
+ mixed_mean = weight * mean + (1 - weight) * mean[partners]
1189
+ mixed_std = weight * std + (1 - weight) * std[partners]
1190
+ return normalized * mixed_std + mixed_mean
1191
+
1192
+
1193
+ class ConvStage(nn.Module):
1194
+ """One convolution, normalization, pooling and dropout stage."""
1195
+
1196
+ def __init__(self, in_channels, out_channels, kernel, dilation, padding, dropout):
1197
+ super().__init__()
1198
+ self.kernel = kernel
1199
+ self.dilation = dilation
1200
+ self.padding = padding
1201
+ self.conv = nn.Conv2d(
1202
+ in_channels,
1203
+ out_channels,
1204
+ kernel,
1205
+ dilation=dilation,
1206
+ padding=padding,
1207
+ bias=False,
1208
+ )
1209
+ self.batch_norm = nn.BatchNorm2d(out_channels)
1210
+ self.pool = nn.AvgPool2d(2)
1211
+ self.dropout = nn.Dropout2d(dropout)
1212
+
1213
+ def forward(self, features):
1214
+ """Transform and downsample one feature map."""
1215
+ features = functional.relu(self.batch_norm(self.conv(features)), inplace=True)
1216
+ return self.dropout(self.pool(features))
1217
+
1218
+ def output_lengths(self, lengths):
1219
+ """Propagate valid time lengths through convolution and pooling."""
1220
+ convolved = lengths + 2 * self.padding[0] - self.dilation[0] * (
1221
+ self.kernel[0] - 1
1222
+ )
1223
+ return torch.div(convolved.clamp_min(0), 2, rounding_mode="floor")
1224
+
1225
+
1226
+ def masked_mean_max(features, lengths):
1227
+ """Pool valid time frames by adding their mean and maximum."""
1228
+ lengths = lengths.clamp(min=1, max=features.size(2))
1229
+ positions = torch.arange(features.size(2), device=features.device).view(1, 1, -1, 1)
1230
+ mask = positions < lengths.view(-1, 1, 1, 1)
1231
+ mean = (features * mask).sum(2) / lengths.view(-1, 1, 1)
1232
+ maximum = features.masked_fill(~mask, float("-inf")).max(2).values
1233
+ return mean + torch.where(torch.isfinite(maximum), maximum, torch.zeros_like(maximum))
1234
+
1235
+
1236
+ class MTRCNNBranch(nn.Module):
1237
+ """One temporal-resolution branch of MTRCNN."""
1238
+
1239
+ def __init__(self, stage_specs, dropout):
1240
+ super().__init__()
1241
+ channels = [1, 16, 32, 64]
1242
+ self.stages = nn.ModuleList(
1243
+ ConvStage(channels[index], channels[index + 1], *specification, dropout)
1244
+ for index, specification in enumerate(stage_specs)
1245
+ )
1246
+ self.mix_style = MixStyle()
1247
+ self.frequency_projection = nn.Linear(self._frequency_bins(), 1)
1248
+
1249
+ def _frequency_bins(self):
1250
+ with torch.no_grad():
1251
+ features = torch.zeros(1, 1, 128, N_MELS)
1252
+ for stage in self.stages:
1253
+ features = stage(features)
1254
+ return features.shape[-1]
1255
+
1256
+ def forward(self, features, lengths, domains, augment):
1257
+ """Encode one temporal-resolution branch."""
1258
+ for index, stage in enumerate(self.stages):
1259
+ features = stage(features)
1260
+ lengths = stage.output_lengths(lengths)
1261
+ if index == 0 and augment:
1262
+ features = self.mix_style(features, domains)
1263
+ pooled = masked_mean_max(features, lengths)
1264
+ return functional.relu(self.frequency_projection(pooled).squeeze(-1), inplace=True)
1265
+
1266
+
1267
+ class GradientReversal(torch.autograd.Function):
1268
+ """Identity forward pass with a sign-reversed backward pass."""
1269
+
1270
+ @staticmethod
1271
+ def forward(ctx, features, strength):
1272
+ """Return features unchanged and retain the reversal strength."""
1273
+ ctx.strength = strength
1274
+ return features.view_as(features)
1275
+
1276
+ @staticmethod
1277
+ def backward(ctx, gradient):
1278
+ """Reverse and scale the upstream embedding gradient."""
1279
+ return -ctx.strength * gradient, None
1280
+
1281
+
1282
+ class MTRCNNDG(nn.Module):
1283
+ """MTRCNN with spectral style, adversarial, and conditional domain heads."""
1284
+
1285
+ def __init__(self, dropout=0.2):
1286
+ super().__init__()
1287
+ self.fourier_mix = FourierMix()
1288
+ self.input_batch_norm = nn.BatchNorm2d(N_MELS)
1289
+ specifications = [
1290
+ [((3, 3), (1, 1), (1, 1)), ((3, 3), (2, 1), (2, 0)), ((3, 3), (3, 1), (3, 0))],
1291
+ [((5, 5), (1, 1), (2, 2)), ((5, 5), (2, 1), (4, 1)), ((5, 5), (3, 1), (6, 1))],
1292
+ [((7, 7), (1, 1), (3, 3)), ((7, 7), (2, 1), (6, 2)), ((7, 7), (3, 1), (9, 2))],
1293
+ ]
1294
+ self.branches = nn.ModuleList(MTRCNNBranch(specs, dropout) for specs in specifications)
1295
+ self.embedding = nn.Linear(64 * 3, 32)
1296
+ self.species_classifier = nn.Linear(32, len(SPECIES))
1297
+ self.domain_classifier = nn.Linear(32, len(DOMAINS))
1298
+ self.aux_domain_classifier = nn.Linear(32, len(DOMAINS))
1299
+
1300
+ def forward(self, features, lengths, domains, domain_strength=0.0, augment=False):
1301
+ """Return species, domain, auxiliary, and embedding outputs."""
1302
+ if augment:
1303
+ features = self.fourier_mix(features, domains)
1304
+ features = features.unsqueeze(1).transpose(1, 3)
1305
+ features = self.input_batch_norm(features).transpose(1, 3)
1306
+ branch_outputs = [
1307
+ branch(features, lengths, domains, augment) for branch in self.branches
1308
+ ]
1309
+ embedding = functional.gelu(self.embedding(torch.cat(branch_outputs, dim=1)))
1310
+ reversed_embedding = GradientReversal.apply(embedding, domain_strength)
1311
+ auxiliary = self.aux_domain_classifier
1312
+ return {
1313
+ "embedding": embedding,
1314
+ "species_logits": self.species_classifier(embedding),
1315
+ "domain_logits": self.domain_classifier(reversed_embedding),
1316
+ "aux_domain_logits": auxiliary(embedding),
1317
+ "frozen_aux_domain_logits": functional.linear(
1318
+ embedding,
1319
+ auxiliary.weight.detach(),
1320
+ auxiliary.bias.detach(),
1321
+ ),
1322
+ }
1323
+
1324
+
1325
+ def supervised_contrastive(embedding, labels, temperature=0.01):
1326
+ """Pull embeddings of the same species together."""
1327
+ normalized = functional.normalize(embedding, dim=1)
1328
+ similarity = normalized @ normalized.T / temperature
1329
+ identity = torch.eye(len(labels), dtype=torch.bool, device=labels.device)
1330
+ positives = labels[:, None].eq(labels[None, :]) & ~identity
1331
+ similarity = similarity.masked_fill(identity, float("-inf"))
1332
+ log_probability = similarity - torch.logsumexp(similarity, dim=1, keepdim=True)
1333
+ log_probability = torch.where(positives, log_probability, torch.zeros_like(log_probability))
1334
+ counts = positives.sum(1)
1335
+ valid = counts > 0
1336
+ if not valid.any():
1337
+ return embedding.sum() * 0
1338
+ return -(log_probability.sum(1)[valid] / counts[valid]).mean()
1339
+
1340
+
1341
+ def conditional_domain_losses(logits, species_labels):
1342
+ """Return class-conditional domain entropy and balance losses."""
1343
+ probabilities = functional.softmax(logits, dim=1).clamp_min(1e-8)
1344
+ entropy_losses = []
1345
+ balance_losses = []
1346
+ for species in species_labels.unique():
1347
+ class_probabilities = probabilities[species_labels == species]
1348
+ entropy_losses.append((class_probabilities * class_probabilities.log()).sum(1).mean())
1349
+ mean_probability = class_probabilities.mean(0)
1350
+ balance_losses.append(
1351
+ (mean_probability * (mean_probability.log() + math.log(len(DOMAINS)))).sum()
1352
+ )
1353
+ return torch.stack(entropy_losses).mean(), torch.stack(balance_losses).mean()
1354
+
1355
+
1356
+ def training_loss(outputs, species_labels, domain_labels, priors, epoch):
1357
+ """Compute the full MTRCNN-DG objective and its components."""
1358
+ adjusted_logits = outputs["species_logits"] + priors.clamp_min(1e-8).log()
1359
+ species_loss = functional.cross_entropy(adjusted_logits, species_labels)
1360
+ domain_loss = functional.cross_entropy(outputs["domain_logits"], domain_labels)
1361
+ auxiliary_loss = functional.cross_entropy(outputs["aux_domain_logits"], domain_labels)
1362
+ contrastive_loss = supervised_contrastive(outputs["embedding"], species_labels)
1363
+ ccde_loss, cdb_loss = conditional_domain_losses(
1364
+ outputs["frozen_aux_domain_logits"], species_labels
1365
+ )
1366
+ conditional_weight = 0.05 if epoch >= 5 else 0.0
1367
+ total = (
1368
+ species_loss
1369
+ + domain_loss
1370
+ + auxiliary_loss
1371
+ + 0.5 * contrastive_loss
1372
+ + conditional_weight * (ccde_loss + cdb_loss)
1373
+ )
1374
+ return {
1375
+ "loss": total,
1376
+ "species": species_loss,
1377
+ "domain": domain_loss,
1378
+ "aux_domain": auxiliary_loss,
1379
+ "contrastive": contrastive_loss,
1380
+ "ccde": ccde_loss,
1381
+ "cdb": cdb_loss,
1382
+ }
1383
+
1384
+
1385
+ def smoke(args):
1386
+ """Exercise frontend, augmentations, all heads, losses, and backward pass."""
1387
+ set_seed(args.seed)
1388
+ multi_domain_species = {0, 1, 2, 6, 7, 8}
1389
+ fold_ids = [
1390
+ f"S_{species + 1}_D_{domain + 1}_{row + 1}"
1391
+ for species in range(len(SPECIES))
1392
+ for domain in range(2 if species in multi_domain_species else 1)
1393
+ for row in range(2)
1394
+ ]
1395
+ training_ids, validation_ids, held_domains = held_domain_split(
1396
+ fold_ids, 0, seed=args.seed
1397
+ )
1398
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
1399
+ dataset = MosquitoDataset(
1400
+ None,
1401
+ synthetic_ids(4),
1402
+ args.seconds,
1403
+ training=True,
1404
+ balance_power=0.5,
1405
+ synthetic=True,
1406
+ )
1407
+ sampler = WeightedRandomSampler(
1408
+ dataset.sample_weights,
1409
+ num_samples=args.batch_size,
1410
+ replacement=True,
1411
+ generator=torch.Generator().manual_seed(args.seed),
1412
+ )
1413
+ batch = next(
1414
+ iter(
1415
+ DataLoader(
1416
+ dataset,
1417
+ batch_size=args.batch_size,
1418
+ sampler=sampler,
1419
+ collate_fn=AudioCollator(),
1420
+ )
1421
+ )
1422
+ )
1423
+ waveforms = batch["waveforms"].to(device)
1424
+ lengths = batch["lengths"].to(device)
1425
+ labels = batch["species"].to(device)
1426
+ domains = batch["domains"].to(device)
1427
+ mean, std = compute_normalization(
1428
+ dataset,
1429
+ size=min(20, len(dataset)),
1430
+ batch_size=args.batch_size,
1431
+ workers=0,
1432
+ device=device,
1433
+ seed=args.seed,
1434
+ min_hz=MIN_HZ,
1435
+ )
1436
+ frontend = LogMelFrontend(mean, std).to(device)
1437
+ model = MTRCNNDG().to(device).train()
1438
+ features, frame_lengths = frontend(waveforms, lengths)
1439
+ outputs = model(features, frame_lengths, domains, domain_strength=0.3, augment=True)
1440
+ priors = torch.full((len(SPECIES),), 1 / len(SPECIES), device=device)
1441
+ losses = training_loss(outputs, labels, domains, priors, epoch=5)
1442
+ losses["loss"].backward()
1443
+ assert outputs["species_logits"].shape == (args.batch_size, len(SPECIES))
1444
+ assert outputs["domain_logits"].shape == (args.batch_size, len(DOMAINS))
1445
+ assert all(torch.isfinite(value) for value in losses.values())
1446
+ assert all(
1447
+ parameter.grad is None or torch.isfinite(parameter.grad).all()
1448
+ for parameter in model.parameters()
1449
+ )
1450
+ result = {
1451
+ "device": str(device),
1452
+ "features": list(features.shape),
1453
+ "embedding": list(outputs["embedding"].shape),
1454
+ "parameters": sum(parameter.numel() for parameter in model.parameters()),
1455
+ "normalization": {
1456
+ "mean": [round(mean.min().item(), 4), round(mean.max().item(), 4)],
1457
+ "std": [round(std.min().item(), 4), round(std.max().item(), 4)],
1458
+ },
1459
+ "fold": {
1460
+ "train": len(training_ids),
1461
+ "validation": len(validation_ids),
1462
+ "held_domains": {
1463
+ SPECIES[species]: DOMAINS[domain] if domain is not None else "same-domain"
1464
+ for species, domain in held_domains.items()
1465
+ },
1466
+ },
1467
+ "losses": {name: round(value.item(), 6) for name, value in losses.items()},
1468
+ }
1469
+ print(json.dumps(result, indent=1))
1470
+
1471
+
1472
+ def add_train_parser(commands):
1473
+ """Register the train command and its options."""
1474
+ train_parser = commands.add_parser("train")
1475
+ train_parser.add_argument("--data", default="/tmp/mosquito-audio")
1476
+ train_parser.add_argument("--output", default="/tmp/mosquito-mtrcnn")
1477
+ train_parser.add_argument("--fold", type=int, default=0)
1478
+ train_parser.add_argument("--epochs", type=int, default=20)
1479
+ train_parser.add_argument("--max-steps", type=int, default=-1)
1480
+ train_parser.add_argument("--train-size", type=int, default=None)
1481
+ train_parser.add_argument("--validation-size", type=int, default=None)
1482
+ train_parser.add_argument("--stats-size", type=int, default=10_000)
1483
+ train_parser.add_argument("--reference-size", type=int, default=20_000)
1484
+ train_parser.add_argument("--batch-size", type=int, default=128)
1485
+ train_parser.add_argument("--eval-batch-size", type=int, default=256)
1486
+ train_parser.add_argument("--learning-rate", type=float, default=1e-3)
1487
+ train_parser.add_argument("--weight-decay", type=float, default=1e-4)
1488
+ train_parser.add_argument("--warmup-ratio", type=float, default=0.05)
1489
+ train_parser.add_argument("--gradient-clip", type=float, default=5.0)
1490
+ train_parser.add_argument("--dropout", type=float, default=0.2)
1491
+ train_parser.add_argument("--balance-power", type=float, default=1.0)
1492
+ train_parser.add_argument("--seconds", type=float, default=SECONDS)
1493
+ train_parser.add_argument("--min-hz", type=float, default=MIN_HZ)
1494
+ train_parser.add_argument("--level-rms", type=float, default=LEVEL_RMS)
1495
+ train_parser.add_argument("--patience", type=int, default=4)
1496
+ train_parser.add_argument("--workers", type=int, default=max(0, cpus() - 1))
1497
+ train_parser.add_argument("--seed", type=int, default=42)
1498
+ train_parser.add_argument("--synthetic", action="store_true")
1499
+ train_parser.add_argument("--synthetic-rows", type=int, default=4)
1500
+ train_parser.add_argument("--test", action="store_true")
1501
+ train_parser.add_argument("--test-size", type=int, default=None)
1502
+ train_parser.set_defaults(run=train)
1503
+
1504
+
1505
+ def main():
1506
+ """Parse command-line arguments."""
1507
+ parser = argparse.ArgumentParser(description=__doc__)
1508
+ commands = parser.add_subparsers(dest="command", required=True)
1509
+ smoke_parser = commands.add_parser("smoke")
1510
+ smoke_parser.add_argument("--batch-size", type=int, default=10)
1511
+ smoke_parser.add_argument("--seconds", type=float, default=SECONDS)
1512
+ smoke_parser.add_argument("--seed", type=int, default=42)
1513
+ smoke_parser.set_defaults(run=smoke)
1514
+ add_train_parser(commands)
1515
+ evaluate_parser = commands.add_parser("evaluate")
1516
+ evaluate_parser.add_argument("--checkpoint", required=True)
1517
+ evaluate_parser.add_argument("--data", default="/tmp/mosquito-audio")
1518
+ evaluate_parser.add_argument("--output", default="/tmp/mosquito-mtrcnn-eval.json")
1519
+ evaluate_parser.add_argument("--test-size", type=int, default=None)
1520
+ evaluate_parser.add_argument("--eval-batch-size", type=int, default=128)
1521
+ evaluate_parser.add_argument("--workers", type=int, default=max(0, cpus() - 1))
1522
+ evaluate_parser.add_argument("--export", default=None)
1523
+ evaluate_parser.add_argument("--repo", default=None)
1524
+ evaluate_parser.set_defaults(run=evaluate)
1525
+ predict_parser = commands.add_parser("predict")
1526
+ predict_parser.add_argument("--model", default=REPO)
1527
+ predict_parser.add_argument("files", nargs="+")
1528
+ predict_parser.set_defaults(run=predict)
1529
+ card_parser = commands.add_parser("card")
1530
+ card_parser.add_argument("--repo", default=REPO)
1531
+ card_parser.set_defaults(run=card)
1532
+ args = parser.parse_args()
1533
+ args.run(args)
1534
+
1535
+
1536
+ if __name__ == "__main__":
1537
+ main()
results/eval.json ADDED
@@ -0,0 +1,1177 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "checkpoint": "/checkpoint/checkpoint.pt",
3
+ "epoch": 12,
4
+ "fold": -1,
5
+ "seed": 42,
6
+ "arguments": {
7
+ "command": "train",
8
+ "data": "/tmp/mosquito-audio",
9
+ "output": "/data/mosquito-mtrcnn-dg-v1-full-seed42",
10
+ "fold": -1,
11
+ "epochs": 12,
12
+ "max_steps": -1,
13
+ "train_size": null,
14
+ "validation_size": null,
15
+ "stats_size": 10000,
16
+ "reference_size": 20000,
17
+ "batch_size": 128,
18
+ "eval_batch_size": 256,
19
+ "learning_rate": 0.001,
20
+ "weight_decay": 0.0001,
21
+ "warmup_ratio": 0.05,
22
+ "gradient_clip": 5.0,
23
+ "dropout": 0.2,
24
+ "balance_power": 1.0,
25
+ "seconds": 2.0,
26
+ "min_hz": 0.0,
27
+ "level_rms": 0.0,
28
+ "patience": 4,
29
+ "workers": 2,
30
+ "seed": 42,
31
+ "synthetic": false,
32
+ "synthetic_rows": 4,
33
+ "test": true,
34
+ "test_size": null
35
+ },
36
+ "test": {
37
+ "rows": 27217,
38
+ "softmax": {
39
+ "seen": {
40
+ "balanced_accuracy": 0.6932280918167799,
41
+ "rows": 23015,
42
+ "per_species": {
43
+ "Aedes aegypti": 0.6460399146479227,
44
+ "Aedes albopictus": 0.44731332868108864,
45
+ "Culex quinquefasciatus": 0.8741965105601469,
46
+ "Anopheles gambiae": 0.7882534775888718,
47
+ "Anopheles arabiensis": 0.6027397260273972,
48
+ "Culex pipiens": 0.8008255933952528
49
+ },
50
+ "confusion": [
51
+ [
52
+ 5147,
53
+ 139,
54
+ 765,
55
+ 598,
56
+ 736,
57
+ 24,
58
+ 555,
59
+ 0,
60
+ 3
61
+ ],
62
+ [
63
+ 82,
64
+ 641,
65
+ 6,
66
+ 533,
67
+ 146,
68
+ 5,
69
+ 19,
70
+ 0,
71
+ 1
72
+ ],
73
+ [
74
+ 146,
75
+ 2,
76
+ 5712,
77
+ 12,
78
+ 623,
79
+ 10,
80
+ 29,
81
+ 0,
82
+ 0
83
+ ],
84
+ [
85
+ 80,
86
+ 28,
87
+ 8,
88
+ 3060,
89
+ 639,
90
+ 4,
91
+ 63,
92
+ 0,
93
+ 0
94
+ ],
95
+ [
96
+ 2,
97
+ 1,
98
+ 13,
99
+ 93,
100
+ 176,
101
+ 0,
102
+ 7,
103
+ 0,
104
+ 0
105
+ ],
106
+ [
107
+ 0,
108
+ 0,
109
+ 0,
110
+ 0,
111
+ 0,
112
+ 0,
113
+ 0,
114
+ 0,
115
+ 0
116
+ ],
117
+ [
118
+ 275,
119
+ 3,
120
+ 49,
121
+ 95,
122
+ 151,
123
+ 4,
124
+ 2328,
125
+ 0,
126
+ 2
127
+ ],
128
+ [
129
+ 0,
130
+ 0,
131
+ 0,
132
+ 0,
133
+ 0,
134
+ 0,
135
+ 0,
136
+ 0,
137
+ 0
138
+ ],
139
+ [
140
+ 0,
141
+ 0,
142
+ 0,
143
+ 0,
144
+ 0,
145
+ 0,
146
+ 0,
147
+ 0,
148
+ 0
149
+ ]
150
+ ]
151
+ },
152
+ "unseen": {
153
+ "balanced_accuracy": 0.2546596781809005,
154
+ "rows": 4202,
155
+ "per_species": {
156
+ "Aedes aegypti": 0.0,
157
+ "Aedes albopictus": 0.29832935560859186,
158
+ "Culex quinquefasciatus": 0.02976190476190476,
159
+ "Anopheles gambiae": 0.0,
160
+ "Anopheles arabiensis": 0.0016483516483516484,
161
+ "Anopheles dirus": 0.0,
162
+ "Culex pipiens": 0.9411764705882353,
163
+ "Anopheles minimus": 0.7777777777777778,
164
+ "Anopheles stephensi": 0.24324324324324326
165
+ },
166
+ "confusion": [
167
+ [
168
+ 0,
169
+ 192,
170
+ 0,
171
+ 0,
172
+ 0,
173
+ 0,
174
+ 0,
175
+ 0,
176
+ 0
177
+ ],
178
+ [
179
+ 81,
180
+ 125,
181
+ 0,
182
+ 9,
183
+ 0,
184
+ 114,
185
+ 28,
186
+ 41,
187
+ 21
188
+ ],
189
+ [
190
+ 121,
191
+ 5,
192
+ 20,
193
+ 0,
194
+ 0,
195
+ 281,
196
+ 6,
197
+ 102,
198
+ 137
199
+ ],
200
+ [
201
+ 61,
202
+ 152,
203
+ 90,
204
+ 0,
205
+ 0,
206
+ 13,
207
+ 0,
208
+ 474,
209
+ 28
210
+ ],
211
+ [
212
+ 544,
213
+ 253,
214
+ 262,
215
+ 0,
216
+ 3,
217
+ 363,
218
+ 12,
219
+ 307,
220
+ 76
221
+ ],
222
+ [
223
+ 1,
224
+ 0,
225
+ 0,
226
+ 0,
227
+ 0,
228
+ 0,
229
+ 0,
230
+ 39,
231
+ 0
232
+ ],
233
+ [
234
+ 2,
235
+ 0,
236
+ 1,
237
+ 0,
238
+ 0,
239
+ 0,
240
+ 64,
241
+ 0,
242
+ 1
243
+ ],
244
+ [
245
+ 0,
246
+ 0,
247
+ 0,
248
+ 0,
249
+ 0,
250
+ 0,
251
+ 0,
252
+ 77,
253
+ 22
254
+ ],
255
+ [
256
+ 0,
257
+ 0,
258
+ 25,
259
+ 0,
260
+ 0,
261
+ 0,
262
+ 0,
263
+ 31,
264
+ 18
265
+ ]
266
+ ]
267
+ }
268
+ },
269
+ "mahalanobis": {
270
+ "seen": {
271
+ "balanced_accuracy": 0.749029147696192,
272
+ "rows": 23015,
273
+ "per_species": {
274
+ "Aedes aegypti": 0.7094263838333125,
275
+ "Aedes albopictus": 0.7997208653175157,
276
+ "Culex quinquefasciatus": 0.8790939700030609,
277
+ "Anopheles gambiae": 0.6254507985574446,
278
+ "Anopheles arabiensis": 0.6335616438356164,
279
+ "Culex pipiens": 0.8469212246302029
280
+ },
281
+ "confusion": [
282
+ [
283
+ 5652,
284
+ 269,
285
+ 630,
286
+ 155,
287
+ 513,
288
+ 1,
289
+ 747,
290
+ 0,
291
+ 0
292
+ ],
293
+ [
294
+ 99,
295
+ 1146,
296
+ 5,
297
+ 75,
298
+ 83,
299
+ 0,
300
+ 25,
301
+ 0,
302
+ 0
303
+ ],
304
+ [
305
+ 268,
306
+ 25,
307
+ 5744,
308
+ 0,
309
+ 443,
310
+ 0,
311
+ 54,
312
+ 0,
313
+ 0
314
+ ],
315
+ [
316
+ 189,
317
+ 317,
318
+ 9,
319
+ 2428,
320
+ 830,
321
+ 0,
322
+ 109,
323
+ 0,
324
+ 0
325
+ ],
326
+ [
327
+ 4,
328
+ 16,
329
+ 18,
330
+ 58,
331
+ 185,
332
+ 0,
333
+ 11,
334
+ 0,
335
+ 0
336
+ ],
337
+ [
338
+ 0,
339
+ 0,
340
+ 0,
341
+ 0,
342
+ 0,
343
+ 0,
344
+ 0,
345
+ 0,
346
+ 0
347
+ ],
348
+ [
349
+ 267,
350
+ 15,
351
+ 40,
352
+ 26,
353
+ 96,
354
+ 0,
355
+ 2462,
356
+ 0,
357
+ 1
358
+ ],
359
+ [
360
+ 0,
361
+ 0,
362
+ 0,
363
+ 0,
364
+ 0,
365
+ 0,
366
+ 0,
367
+ 0,
368
+ 0
369
+ ],
370
+ [
371
+ 0,
372
+ 0,
373
+ 0,
374
+ 0,
375
+ 0,
376
+ 0,
377
+ 0,
378
+ 0,
379
+ 0
380
+ ]
381
+ ]
382
+ },
383
+ "unseen": {
384
+ "balanced_accuracy": 0.2824912329416518,
385
+ "rows": 4202,
386
+ "per_species": {
387
+ "Aedes aegypti": 0.0,
388
+ "Aedes albopictus": 0.5918854415274463,
389
+ "Culex quinquefasciatus": 0.05357142857142857,
390
+ "Anopheles gambiae": 0.0,
391
+ "Anopheles arabiensis": 0.002197802197802198,
392
+ "Anopheles dirus": 0.0,
393
+ "Culex pipiens": 0.9411764705882353,
394
+ "Anopheles minimus": 0.7373737373737373,
395
+ "Anopheles stephensi": 0.21621621621621623
396
+ },
397
+ "confusion": [
398
+ [
399
+ 0,
400
+ 192,
401
+ 0,
402
+ 0,
403
+ 0,
404
+ 0,
405
+ 0,
406
+ 0,
407
+ 0
408
+ ],
409
+ [
410
+ 85,
411
+ 248,
412
+ 0,
413
+ 4,
414
+ 0,
415
+ 53,
416
+ 25,
417
+ 4,
418
+ 0
419
+ ],
420
+ [
421
+ 194,
422
+ 65,
423
+ 36,
424
+ 0,
425
+ 0,
426
+ 302,
427
+ 43,
428
+ 3,
429
+ 29
430
+ ],
431
+ [
432
+ 56,
433
+ 309,
434
+ 221,
435
+ 0,
436
+ 1,
437
+ 15,
438
+ 26,
439
+ 179,
440
+ 11
441
+ ],
442
+ [
443
+ 601,
444
+ 457,
445
+ 330,
446
+ 0,
447
+ 4,
448
+ 327,
449
+ 92,
450
+ 8,
451
+ 1
452
+ ],
453
+ [
454
+ 3,
455
+ 1,
456
+ 2,
457
+ 0,
458
+ 0,
459
+ 0,
460
+ 1,
461
+ 33,
462
+ 0
463
+ ],
464
+ [
465
+ 1,
466
+ 0,
467
+ 2,
468
+ 0,
469
+ 0,
470
+ 0,
471
+ 64,
472
+ 0,
473
+ 1
474
+ ],
475
+ [
476
+ 1,
477
+ 1,
478
+ 1,
479
+ 0,
480
+ 0,
481
+ 0,
482
+ 3,
483
+ 73,
484
+ 20
485
+ ],
486
+ [
487
+ 4,
488
+ 0,
489
+ 25,
490
+ 0,
491
+ 0,
492
+ 0,
493
+ 1,
494
+ 28,
495
+ 16
496
+ ]
497
+ ]
498
+ }
499
+ }
500
+ },
501
+ "fold_softmax": {
502
+ "balanced_accuracy": 0.7731446928816349,
503
+ "rows": 30516,
504
+ "per_species": {
505
+ "Aedes aegypti": 0.6406232973738695,
506
+ "Aedes albopictus": 0.4512722035525684,
507
+ "Culex quinquefasciatus": 0.8714373843306601,
508
+ "Anopheles gambiae": 0.7983733686400605,
509
+ "Anopheles arabiensis": 0.5806315789473684,
510
+ "Anopheles dirus": 1.0,
511
+ "Culex pipiens": 0.8031072602330445,
512
+ "Anopheles minimus": 0.8928571428571429,
513
+ "Anopheles stephensi": 0.92
514
+ },
515
+ "confusion": [
516
+ [
517
+ 5879,
518
+ 136,
519
+ 873,
520
+ 653,
521
+ 928,
522
+ 31,
523
+ 669,
524
+ 1,
525
+ 7
526
+ ],
527
+ [
528
+ 126,
529
+ 940,
530
+ 12,
531
+ 785,
532
+ 176,
533
+ 3,
534
+ 40,
535
+ 0,
536
+ 1
537
+ ],
538
+ [
539
+ 200,
540
+ 2,
541
+ 7063,
542
+ 13,
543
+ 787,
544
+ 13,
545
+ 26,
546
+ 0,
547
+ 1
548
+ ],
549
+ [
550
+ 86,
551
+ 27,
552
+ 11,
553
+ 4221,
554
+ 872,
555
+ 7,
556
+ 63,
557
+ 0,
558
+ 0
559
+ ],
560
+ [
561
+ 41,
562
+ 17,
563
+ 99,
564
+ 746,
565
+ 1379,
566
+ 2,
567
+ 90,
568
+ 0,
569
+ 1
570
+ ],
571
+ [
572
+ 0,
573
+ 0,
574
+ 0,
575
+ 0,
576
+ 0,
577
+ 11,
578
+ 0,
579
+ 0,
580
+ 0
581
+ ],
582
+ [
583
+ 345,
584
+ 1,
585
+ 36,
586
+ 96,
587
+ 179,
588
+ 2,
589
+ 2688,
590
+ 0,
591
+ 0
592
+ ],
593
+ [
594
+ 0,
595
+ 0,
596
+ 0,
597
+ 0,
598
+ 0,
599
+ 3,
600
+ 0,
601
+ 50,
602
+ 3
603
+ ],
604
+ [
605
+ 0,
606
+ 0,
607
+ 0,
608
+ 0,
609
+ 0,
610
+ 0,
611
+ 0,
612
+ 6,
613
+ 69
614
+ ]
615
+ ],
616
+ "domain_holdout": {
617
+ "balanced_accuracy": 0.0,
618
+ "rows": 0,
619
+ "per_species": {},
620
+ "confusion": [
621
+ [
622
+ 0,
623
+ 0,
624
+ 0,
625
+ 0,
626
+ 0,
627
+ 0,
628
+ 0,
629
+ 0,
630
+ 0
631
+ ],
632
+ [
633
+ 0,
634
+ 0,
635
+ 0,
636
+ 0,
637
+ 0,
638
+ 0,
639
+ 0,
640
+ 0,
641
+ 0
642
+ ],
643
+ [
644
+ 0,
645
+ 0,
646
+ 0,
647
+ 0,
648
+ 0,
649
+ 0,
650
+ 0,
651
+ 0,
652
+ 0
653
+ ],
654
+ [
655
+ 0,
656
+ 0,
657
+ 0,
658
+ 0,
659
+ 0,
660
+ 0,
661
+ 0,
662
+ 0,
663
+ 0
664
+ ],
665
+ [
666
+ 0,
667
+ 0,
668
+ 0,
669
+ 0,
670
+ 0,
671
+ 0,
672
+ 0,
673
+ 0,
674
+ 0
675
+ ],
676
+ [
677
+ 0,
678
+ 0,
679
+ 0,
680
+ 0,
681
+ 0,
682
+ 0,
683
+ 0,
684
+ 0,
685
+ 0
686
+ ],
687
+ [
688
+ 0,
689
+ 0,
690
+ 0,
691
+ 0,
692
+ 0,
693
+ 0,
694
+ 0,
695
+ 0,
696
+ 0
697
+ ],
698
+ [
699
+ 0,
700
+ 0,
701
+ 0,
702
+ 0,
703
+ 0,
704
+ 0,
705
+ 0,
706
+ 0,
707
+ 0
708
+ ],
709
+ [
710
+ 0,
711
+ 0,
712
+ 0,
713
+ 0,
714
+ 0,
715
+ 0,
716
+ 0,
717
+ 0,
718
+ 0
719
+ ]
720
+ ]
721
+ },
722
+ "same_domain": {
723
+ "balanced_accuracy": 0.7731446928816349,
724
+ "rows": 30516,
725
+ "per_species": {
726
+ "Aedes aegypti": 0.6406232973738695,
727
+ "Aedes albopictus": 0.4512722035525684,
728
+ "Culex quinquefasciatus": 0.8714373843306601,
729
+ "Anopheles gambiae": 0.7983733686400605,
730
+ "Anopheles arabiensis": 0.5806315789473684,
731
+ "Anopheles dirus": 1.0,
732
+ "Culex pipiens": 0.8031072602330445,
733
+ "Anopheles minimus": 0.8928571428571429,
734
+ "Anopheles stephensi": 0.92
735
+ },
736
+ "confusion": [
737
+ [
738
+ 5879,
739
+ 136,
740
+ 873,
741
+ 653,
742
+ 928,
743
+ 31,
744
+ 669,
745
+ 1,
746
+ 7
747
+ ],
748
+ [
749
+ 126,
750
+ 940,
751
+ 12,
752
+ 785,
753
+ 176,
754
+ 3,
755
+ 40,
756
+ 0,
757
+ 1
758
+ ],
759
+ [
760
+ 200,
761
+ 2,
762
+ 7063,
763
+ 13,
764
+ 787,
765
+ 13,
766
+ 26,
767
+ 0,
768
+ 1
769
+ ],
770
+ [
771
+ 86,
772
+ 27,
773
+ 11,
774
+ 4221,
775
+ 872,
776
+ 7,
777
+ 63,
778
+ 0,
779
+ 0
780
+ ],
781
+ [
782
+ 41,
783
+ 17,
784
+ 99,
785
+ 746,
786
+ 1379,
787
+ 2,
788
+ 90,
789
+ 0,
790
+ 1
791
+ ],
792
+ [
793
+ 0,
794
+ 0,
795
+ 0,
796
+ 0,
797
+ 0,
798
+ 11,
799
+ 0,
800
+ 0,
801
+ 0
802
+ ],
803
+ [
804
+ 345,
805
+ 1,
806
+ 36,
807
+ 96,
808
+ 179,
809
+ 2,
810
+ 2688,
811
+ 0,
812
+ 0
813
+ ],
814
+ [
815
+ 0,
816
+ 0,
817
+ 0,
818
+ 0,
819
+ 0,
820
+ 3,
821
+ 0,
822
+ 50,
823
+ 3
824
+ ],
825
+ [
826
+ 0,
827
+ 0,
828
+ 0,
829
+ 0,
830
+ 0,
831
+ 0,
832
+ 0,
833
+ 6,
834
+ 69
835
+ ]
836
+ ]
837
+ }
838
+ },
839
+ "fold_mahalanobis": {
840
+ "balanced_accuracy": 0.8076237330803099,
841
+ "rows": 30516,
842
+ "per_species": {
843
+ "Aedes aegypti": 0.6983763757219135,
844
+ "Aedes albopictus": 0.7820451272203552,
845
+ "Culex quinquefasciatus": 0.8768661320172733,
846
+ "Anopheles gambiae": 0.6372233780972196,
847
+ "Anopheles arabiensis": 0.6058947368421053,
848
+ "Anopheles dirus": 1.0,
849
+ "Culex pipiens": 0.8598745144905886,
850
+ "Anopheles minimus": 0.875,
851
+ "Anopheles stephensi": 0.9333333333333333
852
+ },
853
+ "confusion": [
854
+ [
855
+ 6409,
856
+ 297,
857
+ 744,
858
+ 165,
859
+ 672,
860
+ 0,
861
+ 890,
862
+ 0,
863
+ 0
864
+ ],
865
+ [
866
+ 146,
867
+ 1629,
868
+ 13,
869
+ 122,
870
+ 120,
871
+ 0,
872
+ 53,
873
+ 0,
874
+ 0
875
+ ],
876
+ [
877
+ 354,
878
+ 28,
879
+ 7107,
880
+ 1,
881
+ 545,
882
+ 0,
883
+ 70,
884
+ 0,
885
+ 0
886
+ ],
887
+ [
888
+ 205,
889
+ 413,
890
+ 13,
891
+ 3369,
892
+ 1160,
893
+ 0,
894
+ 127,
895
+ 0,
896
+ 0
897
+ ],
898
+ [
899
+ 78,
900
+ 125,
901
+ 125,
902
+ 469,
903
+ 1439,
904
+ 0,
905
+ 139,
906
+ 0,
907
+ 0
908
+ ],
909
+ [
910
+ 0,
911
+ 0,
912
+ 0,
913
+ 0,
914
+ 0,
915
+ 11,
916
+ 0,
917
+ 0,
918
+ 0
919
+ ],
920
+ [
921
+ 300,
922
+ 7,
923
+ 37,
924
+ 27,
925
+ 98,
926
+ 0,
927
+ 2878,
928
+ 0,
929
+ 0
930
+ ],
931
+ [
932
+ 0,
933
+ 0,
934
+ 0,
935
+ 0,
936
+ 0,
937
+ 3,
938
+ 1,
939
+ 49,
940
+ 3
941
+ ],
942
+ [
943
+ 0,
944
+ 0,
945
+ 0,
946
+ 0,
947
+ 0,
948
+ 0,
949
+ 0,
950
+ 5,
951
+ 70
952
+ ]
953
+ ],
954
+ "domain_holdout": {
955
+ "balanced_accuracy": 0.0,
956
+ "rows": 0,
957
+ "per_species": {},
958
+ "confusion": [
959
+ [
960
+ 0,
961
+ 0,
962
+ 0,
963
+ 0,
964
+ 0,
965
+ 0,
966
+ 0,
967
+ 0,
968
+ 0
969
+ ],
970
+ [
971
+ 0,
972
+ 0,
973
+ 0,
974
+ 0,
975
+ 0,
976
+ 0,
977
+ 0,
978
+ 0,
979
+ 0
980
+ ],
981
+ [
982
+ 0,
983
+ 0,
984
+ 0,
985
+ 0,
986
+ 0,
987
+ 0,
988
+ 0,
989
+ 0,
990
+ 0
991
+ ],
992
+ [
993
+ 0,
994
+ 0,
995
+ 0,
996
+ 0,
997
+ 0,
998
+ 0,
999
+ 0,
1000
+ 0,
1001
+ 0
1002
+ ],
1003
+ [
1004
+ 0,
1005
+ 0,
1006
+ 0,
1007
+ 0,
1008
+ 0,
1009
+ 0,
1010
+ 0,
1011
+ 0,
1012
+ 0
1013
+ ],
1014
+ [
1015
+ 0,
1016
+ 0,
1017
+ 0,
1018
+ 0,
1019
+ 0,
1020
+ 0,
1021
+ 0,
1022
+ 0,
1023
+ 0
1024
+ ],
1025
+ [
1026
+ 0,
1027
+ 0,
1028
+ 0,
1029
+ 0,
1030
+ 0,
1031
+ 0,
1032
+ 0,
1033
+ 0,
1034
+ 0
1035
+ ],
1036
+ [
1037
+ 0,
1038
+ 0,
1039
+ 0,
1040
+ 0,
1041
+ 0,
1042
+ 0,
1043
+ 0,
1044
+ 0,
1045
+ 0
1046
+ ],
1047
+ [
1048
+ 0,
1049
+ 0,
1050
+ 0,
1051
+ 0,
1052
+ 0,
1053
+ 0,
1054
+ 0,
1055
+ 0,
1056
+ 0
1057
+ ]
1058
+ ]
1059
+ },
1060
+ "same_domain": {
1061
+ "balanced_accuracy": 0.8076237330803099,
1062
+ "rows": 30516,
1063
+ "per_species": {
1064
+ "Aedes aegypti": 0.6983763757219135,
1065
+ "Aedes albopictus": 0.7820451272203552,
1066
+ "Culex quinquefasciatus": 0.8768661320172733,
1067
+ "Anopheles gambiae": 0.6372233780972196,
1068
+ "Anopheles arabiensis": 0.6058947368421053,
1069
+ "Anopheles dirus": 1.0,
1070
+ "Culex pipiens": 0.8598745144905886,
1071
+ "Anopheles minimus": 0.875,
1072
+ "Anopheles stephensi": 0.9333333333333333
1073
+ },
1074
+ "confusion": [
1075
+ [
1076
+ 6409,
1077
+ 297,
1078
+ 744,
1079
+ 165,
1080
+ 672,
1081
+ 0,
1082
+ 890,
1083
+ 0,
1084
+ 0
1085
+ ],
1086
+ [
1087
+ 146,
1088
+ 1629,
1089
+ 13,
1090
+ 122,
1091
+ 120,
1092
+ 0,
1093
+ 53,
1094
+ 0,
1095
+ 0
1096
+ ],
1097
+ [
1098
+ 354,
1099
+ 28,
1100
+ 7107,
1101
+ 1,
1102
+ 545,
1103
+ 0,
1104
+ 70,
1105
+ 0,
1106
+ 0
1107
+ ],
1108
+ [
1109
+ 205,
1110
+ 413,
1111
+ 13,
1112
+ 3369,
1113
+ 1160,
1114
+ 0,
1115
+ 127,
1116
+ 0,
1117
+ 0
1118
+ ],
1119
+ [
1120
+ 78,
1121
+ 125,
1122
+ 125,
1123
+ 469,
1124
+ 1439,
1125
+ 0,
1126
+ 139,
1127
+ 0,
1128
+ 0
1129
+ ],
1130
+ [
1131
+ 0,
1132
+ 0,
1133
+ 0,
1134
+ 0,
1135
+ 0,
1136
+ 11,
1137
+ 0,
1138
+ 0,
1139
+ 0
1140
+ ],
1141
+ [
1142
+ 300,
1143
+ 7,
1144
+ 37,
1145
+ 27,
1146
+ 98,
1147
+ 0,
1148
+ 2878,
1149
+ 0,
1150
+ 0
1151
+ ],
1152
+ [
1153
+ 0,
1154
+ 0,
1155
+ 0,
1156
+ 0,
1157
+ 0,
1158
+ 3,
1159
+ 1,
1160
+ 49,
1161
+ 3
1162
+ ],
1163
+ [
1164
+ 0,
1165
+ 0,
1166
+ 0,
1167
+ 0,
1168
+ 0,
1169
+ 0,
1170
+ 0,
1171
+ 5,
1172
+ 70
1173
+ ]
1174
+ ]
1175
+ }
1176
+ }
1177
+ }