74 lines
2.6 KiB
Python
74 lines
2.6 KiB
Python
|
|
import argparse
|
||
|
|
import json
|
||
|
|
from collections import Counter, defaultdict
|
||
|
|
from typing import Any
|
||
|
|
|
||
|
|
import numpy as np
|
||
|
|
|
||
|
|
from app.audio import classify_sequence, load_reference_sequences
|
||
|
|
from app.profile_store import profile_store
|
||
|
|
|
||
|
|
|
||
|
|
def validate_profile(profile_code: str) -> dict[str, Any]:
|
||
|
|
profile_store.reload()
|
||
|
|
profile = profile_store.get(profile_code, "AUDIO_CLASSIFICATION")
|
||
|
|
sequences, labels, files = load_reference_sequences(profile)
|
||
|
|
confusion: dict[str, Counter[str]] = defaultdict(Counter)
|
||
|
|
details: list[dict[str, Any]] = []
|
||
|
|
correct = 0
|
||
|
|
rejected = 0
|
||
|
|
for index, (query, expected, file_name) in enumerate(zip(sequences, labels, files)):
|
||
|
|
other_indices = [candidate for candidate in range(len(sequences)) if candidate != index]
|
||
|
|
result = classify_sequence(
|
||
|
|
profile,
|
||
|
|
query,
|
||
|
|
[sequences[candidate] for candidate in other_indices],
|
||
|
|
labels[other_indices],
|
||
|
|
files[other_indices],
|
||
|
|
)
|
||
|
|
predicted = result["label"]
|
||
|
|
if predicted == profile.config.get("decision", {}).get("unknownLabel", "UNKNOWN"):
|
||
|
|
rejected += 1
|
||
|
|
if predicted == expected:
|
||
|
|
correct += 1
|
||
|
|
confusion[expected][predicted] += 1
|
||
|
|
ranking = sorted(result["scores"].items(), key=lambda item: item[1], reverse=True)
|
||
|
|
margin = ranking[0][1] - ranking[1][1] if len(ranking) > 1 else ranking[0][1]
|
||
|
|
details.append(
|
||
|
|
{
|
||
|
|
"file": file_name,
|
||
|
|
"expected": expected,
|
||
|
|
"predicted": predicted,
|
||
|
|
"confidence": result["confidence"],
|
||
|
|
"scoreMargin": round(margin, 6),
|
||
|
|
}
|
||
|
|
)
|
||
|
|
|
||
|
|
total = len(details)
|
||
|
|
return {
|
||
|
|
"profileCode": profile.code,
|
||
|
|
"method": "leave-one-out",
|
||
|
|
"total": total,
|
||
|
|
"correct": correct,
|
||
|
|
"accuracy": round(correct / total, 6) if total else 0.0,
|
||
|
|
"rejected": rejected,
|
||
|
|
"thresholds": profile.config.get("decision", {}),
|
||
|
|
"confusion": {label: dict(values) for label, values in sorted(confusion.items())},
|
||
|
|
"errors": [item for item in details if item["expected"] != item["predicted"]],
|
||
|
|
}
|
||
|
|
|
||
|
|
|
||
|
|
def main() -> None:
|
||
|
|
parser = argparse.ArgumentParser(description="Validate an audio feature library without self-matching")
|
||
|
|
parser.add_argument("--profile-code", required=True)
|
||
|
|
parser.add_argument("--fail-below", type=float, default=0.0)
|
||
|
|
args = parser.parse_args()
|
||
|
|
result = validate_profile(args.profile_code)
|
||
|
|
print(json.dumps(result, ensure_ascii=False, indent=2))
|
||
|
|
if result["accuracy"] < args.fail_below:
|
||
|
|
raise SystemExit(2)
|
||
|
|
|
||
|
|
|
||
|
|
if __name__ == "__main__":
|
||
|
|
main()
|