242 lines
9.7 KiB
Python
242 lines
9.7 KiB
Python
import json
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
import numpy as np
|
|
from scipy import signal
|
|
from scipy.fft import dct
|
|
from scipy.spatial.distance import cdist
|
|
|
|
from app.media import decode_audio
|
|
from app.profile_store import Profile, artifact_path
|
|
|
|
|
|
SUPPORTED_AUDIO_SUFFIXES = {".wav", ".mp3", ".m4a", ".aac", ".flac", ".ogg"}
|
|
|
|
|
|
def _hz_to_mel(hz: np.ndarray | float) -> np.ndarray | float:
|
|
return 2595.0 * np.log10(1.0 + np.asarray(hz) / 700.0)
|
|
|
|
|
|
def _mel_to_hz(mel: np.ndarray) -> np.ndarray:
|
|
return 700.0 * (np.power(10.0, mel / 2595.0) - 1.0)
|
|
|
|
|
|
def _mel_filterbank(sample_rate: int, n_fft: int, bands: int) -> np.ndarray:
|
|
mel_points = np.linspace(_hz_to_mel(40.0), _hz_to_mel(sample_rate / 2), bands + 2)
|
|
bins = np.floor((n_fft + 1) * _mel_to_hz(mel_points) / sample_rate).astype(int)
|
|
bins = np.clip(bins, 0, n_fft // 2)
|
|
filters = np.zeros((bands, n_fft // 2 + 1), dtype=np.float32)
|
|
for index in range(bands):
|
|
left, center, right = bins[index : index + 3]
|
|
if center > left:
|
|
filters[index, left:center] = np.linspace(0, 1, center - left, endpoint=False)
|
|
if right > center:
|
|
filters[index, center:right] = np.linspace(1, 0, right - center, endpoint=False)
|
|
return filters
|
|
|
|
|
|
def _trim_audio(samples: np.ndarray, sample_rate: int) -> np.ndarray:
|
|
if len(samples) < sample_rate // 4:
|
|
return samples
|
|
frame = max(1, int(sample_rate * 0.02))
|
|
energy = np.convolve(samples * samples, np.ones(frame) / frame, mode="same")
|
|
threshold = max(float(np.max(energy)) * 0.0025, 1e-7)
|
|
active = np.flatnonzero(energy >= threshold)
|
|
if not len(active):
|
|
return samples
|
|
padding = int(sample_rate * 0.15)
|
|
return samples[max(0, active[0] - padding) : min(len(samples), active[-1] + padding)]
|
|
|
|
|
|
def feature_sequence(
|
|
samples: np.ndarray, sample_rate: int, max_frames: int = 320
|
|
) -> np.ndarray:
|
|
samples = _trim_audio(samples.astype(np.float32), sample_rate)
|
|
samples -= float(np.mean(samples))
|
|
peak = float(np.max(np.abs(samples))) if len(samples) else 0.0
|
|
if peak < 1e-5:
|
|
raise ValueError("Audio contains no usable sound")
|
|
samples /= peak
|
|
|
|
n_fft = 512
|
|
if len(samples) < n_fft:
|
|
samples = np.pad(samples, (0, n_fft - len(samples)))
|
|
_, _, stft = signal.stft(
|
|
samples,
|
|
fs=sample_rate,
|
|
window="hann",
|
|
nperseg=n_fft,
|
|
noverlap=384,
|
|
boundary=None,
|
|
padded=False,
|
|
)
|
|
mel = _mel_filterbank(sample_rate, n_fft, 40) @ (np.abs(stft) ** 2)
|
|
coefficients = dct(np.log(mel + 1e-8), axis=0, norm="ortho")[1:21].T
|
|
|
|
mean = np.mean(coefficients, axis=0)
|
|
static = mean / max(float(np.linalg.norm(mean)), 1e-5)
|
|
normalized = (coefficients - mean) / np.maximum(np.std(coefficients, axis=0), 1e-5)
|
|
delta = np.gradient(normalized, axis=0)
|
|
sequence = np.concatenate(
|
|
[normalized, delta, np.tile(static * 0.75, (len(normalized), 1))], axis=1
|
|
).astype(np.float32)
|
|
sequence /= np.maximum(np.linalg.norm(sequence, axis=1, keepdims=True), 1e-5)
|
|
sequence = sequence[::2]
|
|
if len(sequence) > max_frames:
|
|
indices = np.linspace(0, len(sequence) - 1, max_frames).round().astype(int)
|
|
sequence = sequence[indices]
|
|
return sequence
|
|
|
|
|
|
def sequence_distance(first: np.ndarray, second: np.ndarray, band_ratio: float = 0.35) -> float:
|
|
distances = cdist(first, second, metric="cosine")
|
|
first_length, second_length = distances.shape
|
|
band = max(
|
|
abs(first_length - second_length) + 2,
|
|
int(max(first_length, second_length) * band_ratio),
|
|
)
|
|
previous = np.full(second_length + 1, np.inf, dtype=np.float32)
|
|
previous[0] = 0.0
|
|
for first_index in range(1, first_length + 1):
|
|
current = np.full(second_length + 1, np.inf, dtype=np.float32)
|
|
start = max(1, first_index - band)
|
|
end = min(second_length, first_index + band)
|
|
for second_index in range(start, end + 1):
|
|
current[second_index] = distances[first_index - 1, second_index - 1] + min(
|
|
current[second_index - 1],
|
|
previous[second_index],
|
|
previous[second_index - 1],
|
|
)
|
|
previous = current
|
|
return float(previous[second_length] / (first_length + second_length))
|
|
|
|
|
|
def _read_audio(path: Path) -> tuple[np.ndarray, int]:
|
|
raw, sample_rate = decode_audio(path)
|
|
return np.frombuffer(raw, dtype="<i2").astype(np.float32) / 32768.0, sample_rate
|
|
|
|
|
|
def build_audio_profile(profile: Profile) -> dict[str, Any]:
|
|
references = profile.directory / "references"
|
|
labels: list[str] = []
|
|
files: list[str] = []
|
|
durations: list[float] = []
|
|
sequences: list[np.ndarray] = []
|
|
for label_dir in sorted(path for path in references.glob("*") if path.is_dir()):
|
|
for path in sorted(label_dir.iterdir()):
|
|
if path.suffix.lower() not in SUPPORTED_AUDIO_SUFFIXES:
|
|
continue
|
|
samples, sample_rate = _read_audio(path)
|
|
labels.append(label_dir.name)
|
|
files.append(path.name)
|
|
durations.append(len(samples) / sample_rate)
|
|
sequences.append(feature_sequence(samples, sample_rate))
|
|
|
|
unique_labels = sorted(set(labels))
|
|
if len(unique_labels) < 2:
|
|
raise ValueError("At least two non-empty reference audio groups are required")
|
|
minimum_per_label = int(profile.config.get("minimumReferencesPerLabel", 1))
|
|
counts = {label: labels.count(label) for label in unique_labels}
|
|
insufficient = [label for label, count in counts.items() if count < minimum_per_label]
|
|
if insufficient:
|
|
raise ValueError(
|
|
f"Reference groups below minimum {minimum_per_label}: {', '.join(insufficient)}"
|
|
)
|
|
|
|
offsets = np.concatenate(([0], np.cumsum([len(sequence) for sequence in sequences])))
|
|
output = artifact_path(profile)
|
|
output.parent.mkdir(parents=True, exist_ok=True)
|
|
np.savez_compressed(
|
|
output,
|
|
sequence_features=np.concatenate(sequences),
|
|
sequence_offsets=offsets,
|
|
labels=np.asarray(labels),
|
|
files=np.asarray(files),
|
|
durations=np.asarray(durations, dtype=np.float32),
|
|
)
|
|
manifest = output.with_name("manifest.json")
|
|
manifest.write_text(
|
|
json.dumps(
|
|
{
|
|
"profileCode": profile.code,
|
|
"featureVersion": "mfcc-dtw-v2",
|
|
"labelCounts": counts,
|
|
"referenceCount": len(sequences),
|
|
"files": files,
|
|
},
|
|
ensure_ascii=False,
|
|
indent=2,
|
|
),
|
|
encoding="utf-8",
|
|
)
|
|
return {"referenceCount": len(sequences), "labelCounts": counts}
|
|
|
|
|
|
def load_reference_sequences(profile: Profile):
|
|
artifact = artifact_path(profile)
|
|
if not artifact.exists():
|
|
raise ValueError(
|
|
f"Profile {profile.code} has no feature library; add reference audio and rebuild it"
|
|
)
|
|
data = np.load(artifact, allow_pickle=False)
|
|
if "sequence_features" not in data or "sequence_offsets" not in data:
|
|
raise ValueError(f"Profile {profile.code} uses an obsolete feature library; rebuild it")
|
|
features = data["sequence_features"]
|
|
offsets = data["sequence_offsets"].astype(int)
|
|
sequences = [features[offsets[index] : offsets[index + 1]] for index in range(len(offsets) - 1)]
|
|
return sequences, data["labels"].astype(str), data["files"].astype(str)
|
|
|
|
|
|
def classify_sequence(
|
|
profile: Profile,
|
|
query: np.ndarray,
|
|
reference_sequences: list[np.ndarray],
|
|
labels: np.ndarray,
|
|
files: np.ndarray,
|
|
) -> dict[str, Any]:
|
|
best_by_label: dict[str, dict[str, Any]] = {}
|
|
for reference, label, file_name in zip(reference_sequences, labels, files):
|
|
distance = sequence_distance(query, reference)
|
|
current = best_by_label.get(label)
|
|
if current is None or distance < current["distance"]:
|
|
best_by_label[label] = {"distance": distance, "reference": file_name}
|
|
|
|
ranking = sorted(best_by_label.items(), key=lambda item: item[1]["distance"])
|
|
winner_label, winner = ranking[0]
|
|
second_distance = ranking[1][1]["distance"] if len(ranking) > 1 else 1.0
|
|
distance_margin = second_distance - winner["distance"]
|
|
decision = profile.config.get("decision", {})
|
|
maximum_distance = float(decision.get("maxDistance", 0.45))
|
|
minimum_margin = float(decision.get("minDistanceMargin", 0.015))
|
|
matched = winner["distance"] <= maximum_distance and distance_margin >= minimum_margin
|
|
output_label = winner_label if matched else str(decision.get("unknownLabel", "UNKNOWN"))
|
|
label_names = profile.config.get("labelNames", {})
|
|
return {
|
|
"label": output_label,
|
|
"labelName": label_names.get(output_label, output_label),
|
|
"matched": matched,
|
|
"confidence": round(max(0.0, min(1.0, 1.0 - winner["distance"])), 6),
|
|
"scores": {
|
|
label: round(max(0.0, min(1.0, 1.0 - item["distance"])), 6)
|
|
for label, item in ranking
|
|
},
|
|
"evidence": {"reference": winner["reference"] if matched else None},
|
|
"thresholds": {
|
|
"maxDistance": maximum_distance,
|
|
"minDistanceMargin": minimum_margin,
|
|
},
|
|
}
|
|
|
|
|
|
def classify_audio(profile: Profile, media_path: Path) -> dict[str, Any]:
|
|
reference_sequences, labels, files = load_reference_sequences(profile)
|
|
samples, sample_rate = _read_audio(media_path)
|
|
maximum_seconds = float(profile.config.get("maximumAudioSeconds", 30))
|
|
if len(samples) / sample_rate > maximum_seconds:
|
|
raise ValueError(f"Audio duration exceeds the profile limit of {maximum_seconds:g} seconds")
|
|
query = feature_sequence(samples, sample_rate)
|
|
result = classify_sequence(profile, query, reference_sequences, labels, files)
|
|
result["evidence"].update({"startMs": 0, "endMs": round(len(samples) * 1000 / sample_rate)})
|
|
return result
|