478 lines
16 KiB
Python
478 lines
16 KiB
Python
#!/usr/bin/env python3
|
|
"""Python verifier for Tom game's tiny keyword spotting model.
|
|
|
|
This script mirrors src/recognition/TinyKwsRecognizer.cpp closely enough for
|
|
Windows-side validation without building TensorFlow Lite C++ on Windows.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import math
|
|
import sys
|
|
import wave
|
|
from pathlib import Path
|
|
from typing import Iterable, Tuple
|
|
|
|
import numpy as np
|
|
|
|
|
|
REPO_ROOT = Path(__file__).resolve().parents[4]
|
|
DEFAULT_MODEL = (
|
|
REPO_ROOT
|
|
/ "src"
|
|
/ "Apps"
|
|
/ "Game"
|
|
/ "model"
|
|
/ "mlcommons-tiny-kws"
|
|
/ "models"
|
|
/ "kws_ref_model_float32.tflite"
|
|
)
|
|
|
|
LABELS = [
|
|
"Down",
|
|
"Go",
|
|
"Left",
|
|
"No",
|
|
"Off",
|
|
"On",
|
|
"Right",
|
|
"Stop",
|
|
"Up",
|
|
"Yes",
|
|
"Silence",
|
|
"Unknown",
|
|
]
|
|
|
|
COMMAND_LABELS = {"Down", "Go", "Left", "No", "On", "Right", "Stop", "Up", "Yes"}
|
|
|
|
TARGET_SAMPLE_RATE = 16000
|
|
TARGET_CHANNELS = 1
|
|
WINDOW_SAMPLES = 16000
|
|
WINDOW_STEP_SAMPLES = 8000
|
|
FFT_SIZE = 512
|
|
SPECTRUM_BINS = FFT_SIZE // 2 + 1
|
|
FRAME_STEP_SAMPLES = 320
|
|
FEATURE_FRAME_COUNT = 49
|
|
MEL_BIN_COUNT = 40
|
|
MFCC_BIN_COUNT = 10
|
|
|
|
|
|
def load_interpreter(model_path: Path):
|
|
try:
|
|
from tflite_runtime.interpreter import Interpreter
|
|
|
|
backend = "tflite_runtime"
|
|
except ImportError:
|
|
try:
|
|
import tensorflow as tf
|
|
|
|
Interpreter = tf.lite.Interpreter
|
|
backend = "tensorflow.lite"
|
|
except Exception as exc:
|
|
raise RuntimeError(
|
|
"No TFLite interpreter backend is available. Install either "
|
|
"`tflite-runtime` or `tensorflow` in the active Conda env."
|
|
) from exc
|
|
|
|
interpreter = Interpreter(model_path=str(model_path), num_threads=1)
|
|
interpreter.allocate_tensors()
|
|
return backend, interpreter
|
|
|
|
|
|
def read_wav_int16(path: Path) -> Tuple[np.ndarray, int, int]:
|
|
with wave.open(str(path), "rb") as wav:
|
|
sample_width = wav.getsampwidth()
|
|
if sample_width != 2:
|
|
raise ValueError(f"{path} is {sample_width * 8}-bit PCM; expected 16-bit PCM")
|
|
|
|
sample_rate = wav.getframerate()
|
|
channels = wav.getnchannels()
|
|
frame_count = wav.getnframes()
|
|
raw = wav.readframes(frame_count)
|
|
|
|
samples = np.frombuffer(raw, dtype="<i2").astype(np.int16, copy=False)
|
|
return samples, sample_rate, channels
|
|
|
|
|
|
def write_wav_int16(path: Path, samples: np.ndarray, sample_rate: int, channels: int) -> None:
|
|
path.parent.mkdir(parents=True, exist_ok=True)
|
|
with wave.open(str(path), "wb") as wav:
|
|
wav.setnchannels(channels)
|
|
wav.setsampwidth(2)
|
|
wav.setframerate(sample_rate)
|
|
wav.writeframes(np.asarray(samples, dtype="<i2").tobytes())
|
|
|
|
|
|
def list_devices() -> None:
|
|
try:
|
|
import sounddevice as sd
|
|
except ImportError as exc:
|
|
raise RuntimeError("`sounddevice` is required for --list-devices") from exc
|
|
|
|
print(sd.query_devices())
|
|
|
|
|
|
def record_audio(seconds: float, sample_rate: int, channels: int, device: str | None) -> np.ndarray:
|
|
try:
|
|
import sounddevice as sd
|
|
except ImportError as exc:
|
|
raise RuntimeError("`sounddevice` is required for --record") from exc
|
|
|
|
frame_count = int(round(seconds * sample_rate))
|
|
print(
|
|
f"[INFO] Recording {seconds:.2f}s at {sample_rate} Hz, "
|
|
f"{channels} channel(s), device={device if device is not None else 'default'}"
|
|
)
|
|
audio = sd.rec(
|
|
frame_count,
|
|
samplerate=sample_rate,
|
|
channels=channels,
|
|
dtype="int16",
|
|
device=device,
|
|
)
|
|
sd.wait()
|
|
return np.asarray(audio, dtype=np.int16).reshape(-1)
|
|
|
|
|
|
def convert_channels(samples: np.ndarray, source_channels: int, target_channels: int) -> np.ndarray:
|
|
if samples.size == 0 or source_channels == target_channels:
|
|
return samples.astype(np.int16, copy=False)
|
|
|
|
frames = samples.reshape((-1, source_channels))
|
|
if target_channels == 1:
|
|
mixed = np.mean(frames.astype(np.int32), axis=1)
|
|
return np.clip(np.round(mixed), -32768, 32767).astype(np.int16)
|
|
|
|
output = np.zeros((frames.shape[0], target_channels), dtype=np.int16)
|
|
for channel in range(target_channels):
|
|
output[:, channel] = frames[:, min(channel, source_channels - 1)]
|
|
return output.reshape(-1)
|
|
|
|
|
|
def resample_linear(
|
|
samples: np.ndarray,
|
|
source_rate: int,
|
|
target_rate: int,
|
|
channels: int,
|
|
) -> np.ndarray:
|
|
if samples.size == 0 or source_rate == target_rate:
|
|
return samples.astype(np.int16, copy=False)
|
|
|
|
frames = samples.reshape((-1, channels)).astype(np.float32)
|
|
source_frame_count = frames.shape[0]
|
|
target_frame_count = max(1, int(math.ceil(source_frame_count * target_rate / source_rate)))
|
|
|
|
source_positions = np.arange(target_frame_count, dtype=np.float32) * (source_rate / target_rate)
|
|
source0 = np.minimum(source_positions.astype(np.int64), source_frame_count - 1)
|
|
source1 = np.minimum(source0 + 1, source_frame_count - 1)
|
|
t = (source_positions - source0).reshape((-1, 1))
|
|
|
|
mixed = frames[source0] + (frames[source1] - frames[source0]) * t
|
|
return np.clip(mixed, -32768, 32767).astype(np.int16).reshape(-1)
|
|
|
|
|
|
def trim_silence(samples: np.ndarray, threshold: float, channels: int) -> np.ndarray:
|
|
if samples.size == 0:
|
|
return samples
|
|
|
|
frames = samples.reshape((-1, channels))
|
|
loud = np.any(np.abs(frames.astype(np.float32) / 32768.0) >= threshold, axis=1)
|
|
indices = np.flatnonzero(loud)
|
|
if indices.size == 0:
|
|
return samples
|
|
|
|
start = int(indices[0])
|
|
end = int(indices[-1]) + 1
|
|
return frames[start:end].reshape(-1).astype(np.int16, copy=False)
|
|
|
|
|
|
def prepare_audio(samples: np.ndarray, sample_rate: int, channels: int, trim: bool = True) -> np.ndarray:
|
|
if samples.size == 0 or sample_rate <= 0 or channels <= 0:
|
|
return np.array([], dtype=np.int16)
|
|
|
|
prepared = samples.astype(np.int16, copy=False)
|
|
prepared_channels = channels
|
|
|
|
if prepared_channels != TARGET_CHANNELS:
|
|
prepared = convert_channels(prepared, prepared_channels, TARGET_CHANNELS)
|
|
prepared_channels = TARGET_CHANNELS
|
|
|
|
if sample_rate != TARGET_SAMPLE_RATE:
|
|
prepared = resample_linear(prepared, sample_rate, TARGET_SAMPLE_RATE, prepared_channels)
|
|
|
|
if not trim:
|
|
return prepared
|
|
|
|
trimmed = trim_silence(prepared, 0.02, TARGET_CHANNELS)
|
|
return trimmed if trimmed.size > 0 else prepared
|
|
|
|
|
|
def hertz_to_mel(hertz: np.ndarray | float) -> np.ndarray | float:
|
|
return 2595.0 * np.log10(1.0 + np.asarray(hertz) / 700.0)
|
|
|
|
|
|
def mel_to_hertz(mel: np.ndarray | float) -> np.ndarray | float:
|
|
return 700.0 * (np.power(10.0, np.asarray(mel) / 2595.0) - 1.0)
|
|
|
|
|
|
def make_mel_filter_bank() -> np.ndarray:
|
|
filters = np.zeros((MEL_BIN_COUNT, SPECTRUM_BINS), dtype=np.float32)
|
|
min_mel = float(hertz_to_mel(20.0))
|
|
max_mel = float(hertz_to_mel(4000.0))
|
|
mel_points = np.linspace(min_mel, max_mel, MEL_BIN_COUNT + 2, dtype=np.float32)
|
|
hertz_points = mel_to_hertz(mel_points).astype(np.float32)
|
|
frequencies = (
|
|
np.arange(SPECTRUM_BINS, dtype=np.float32) * TARGET_SAMPLE_RATE / float(FFT_SIZE)
|
|
)
|
|
|
|
for mel in range(MEL_BIN_COUNT):
|
|
left = hertz_points[mel]
|
|
center = hertz_points[mel + 1]
|
|
right = hertz_points[mel + 2]
|
|
|
|
if center > left:
|
|
rising = (frequencies >= left) & (frequencies <= center)
|
|
filters[mel, rising] = (frequencies[rising] - left) / (center - left)
|
|
|
|
if right > center:
|
|
falling = (frequencies > center) & (frequencies <= right)
|
|
filters[mel, falling] = (right - frequencies[falling]) / (right - center)
|
|
|
|
return filters
|
|
|
|
|
|
HANN_WINDOW = (0.5 - 0.5 * np.cos((2.0 * np.pi * np.arange(FFT_SIZE)) / (FFT_SIZE - 1))).astype(
|
|
np.float32
|
|
)
|
|
MEL_FILTER_BANK = make_mel_filter_bank()
|
|
DCT_BASIS = np.zeros((MFCC_BIN_COUNT, MEL_BIN_COUNT), dtype=np.float32)
|
|
for coeff in range(MFCC_BIN_COUNT):
|
|
norm = math.sqrt(1.0 / MEL_BIN_COUNT) if coeff == 0 else math.sqrt(2.0 / MEL_BIN_COUNT)
|
|
for mel in range(MEL_BIN_COUNT):
|
|
DCT_BASIS[coeff, mel] = norm * math.cos(
|
|
math.pi * (mel + 0.5) * coeff / float(MEL_BIN_COUNT)
|
|
)
|
|
|
|
|
|
def normalize_window(samples: np.ndarray, start: int) -> np.ndarray:
|
|
output = np.zeros(WINDOW_SAMPLES, dtype=np.float32)
|
|
available = max(0, min(WINDOW_SAMPLES, samples.size - start))
|
|
if available <= 0:
|
|
return output
|
|
|
|
window = samples[start : start + available].astype(np.float32)
|
|
max_abs = float(np.max(np.abs(window)))
|
|
if max_abs <= 0.0:
|
|
return output
|
|
|
|
output[:available] = window / max_abs
|
|
return output
|
|
|
|
|
|
def extract_mfcc_features(window_samples: np.ndarray) -> np.ndarray:
|
|
features = np.zeros((FEATURE_FRAME_COUNT, MFCC_BIN_COUNT), dtype=np.float32)
|
|
|
|
for frame in range(FEATURE_FRAME_COUNT):
|
|
offset = frame * FRAME_STEP_SAMPLES
|
|
frame_samples = np.zeros(FFT_SIZE, dtype=np.float32)
|
|
available = max(0, min(FFT_SIZE, window_samples.size - offset))
|
|
if available > 0:
|
|
frame_samples[:available] = window_samples[offset : offset + available]
|
|
|
|
spectrum = np.abs(np.fft.rfft(frame_samples * HANN_WINDOW, n=FFT_SIZE)).astype(np.float32)
|
|
mel_energy = MEL_FILTER_BANK @ spectrum
|
|
log_mel = np.log(mel_energy + 1.0e-6).astype(np.float32)
|
|
features[frame, :] = DCT_BASIS @ log_mel
|
|
|
|
return features
|
|
|
|
|
|
def tensor_from_features(features: np.ndarray, input_details: dict) -> np.ndarray:
|
|
input_shape = tuple(int(v) for v in input_details["shape"])
|
|
tensor = features.astype(np.float32).reshape(input_shape)
|
|
dtype = input_details["dtype"]
|
|
|
|
if dtype == np.float32:
|
|
return tensor
|
|
|
|
scale, zero_point = input_details.get("quantization", (0.0, 0))
|
|
if scale == 0:
|
|
raise ValueError(f"Unsupported quantized input without scale: {input_details}")
|
|
|
|
quantized = np.round(tensor / scale + zero_point)
|
|
info = np.iinfo(dtype)
|
|
return np.clip(quantized, info.min, info.max).astype(dtype)
|
|
|
|
|
|
def dequantize_output(output: np.ndarray, output_details: dict) -> np.ndarray:
|
|
if output.dtype == np.float32:
|
|
return output.astype(np.float32).reshape(-1)
|
|
|
|
scale, zero_point = output_details.get("quantization", (0.0, 0))
|
|
if scale == 0:
|
|
return output.astype(np.float32).reshape(-1)
|
|
|
|
return ((output.astype(np.float32) - zero_point) * scale).reshape(-1)
|
|
|
|
|
|
def run_window(interpreter, window_samples: np.ndarray) -> np.ndarray:
|
|
input_details = interpreter.get_input_details()[0]
|
|
output_details = interpreter.get_output_details()[0]
|
|
features = extract_mfcc_features(window_samples)
|
|
tensor = tensor_from_features(features, input_details)
|
|
|
|
interpreter.set_tensor(input_details["index"], tensor)
|
|
interpreter.invoke()
|
|
raw_output = interpreter.get_tensor(output_details["index"])
|
|
return dequantize_output(raw_output, output_details)
|
|
|
|
|
|
def window_starts(sample_count: int) -> Iterable[int]:
|
|
if sample_count <= WINDOW_SAMPLES:
|
|
yield 0
|
|
return
|
|
|
|
start = 0
|
|
while start + WINDOW_SAMPLES <= sample_count:
|
|
yield start
|
|
start += WINDOW_STEP_SAMPLES
|
|
|
|
last_start = sample_count - WINDOW_SAMPLES
|
|
if last_start % WINDOW_STEP_SAMPLES != 0:
|
|
yield last_start
|
|
|
|
|
|
def recognize(
|
|
interpreter,
|
|
samples: np.ndarray,
|
|
sample_rate: int,
|
|
channels: int,
|
|
threshold: float,
|
|
trim: bool,
|
|
) -> Tuple[str, float, np.ndarray, int]:
|
|
prepared = prepare_audio(samples, sample_rate, channels, trim=trim)
|
|
if prepared.size == 0:
|
|
raise ValueError("No audio samples available after preprocessing")
|
|
|
|
best_command_label = "None"
|
|
best_command_confidence = 0.0
|
|
best_command_probs = np.zeros(len(LABELS), dtype=np.float32)
|
|
best_command_start = 0
|
|
|
|
# 用于显示:取所有滑窗中“最高类置信度”最大的那一窗,
|
|
# 这样即使最高类不是游戏命令,print_topk 也不会显示全 0。
|
|
best_overall_probs = np.zeros(len(LABELS), dtype=np.float32)
|
|
best_overall_max_confidence = 0.0
|
|
|
|
for start in window_starts(prepared.size):
|
|
probs = run_window(interpreter, normalize_window(prepared, start))
|
|
if probs.size < len(LABELS):
|
|
raise ValueError(f"Model returned {probs.size} values; expected at least {len(LABELS)}")
|
|
|
|
probs_slice = probs[: len(LABELS)]
|
|
class_index = int(np.argmax(probs_slice))
|
|
label = LABELS[class_index]
|
|
confidence = float(probs_slice[class_index])
|
|
is_command = label in COMMAND_LABELS
|
|
|
|
if confidence > best_overall_max_confidence:
|
|
best_overall_max_confidence = confidence
|
|
best_overall_probs = probs_slice.copy()
|
|
|
|
if is_command and confidence > best_command_confidence:
|
|
best_command_label = label
|
|
best_command_confidence = confidence
|
|
best_command_probs = probs_slice.copy()
|
|
best_command_start = start
|
|
|
|
if best_command_confidence < threshold:
|
|
return "None", best_command_confidence, best_overall_probs, best_command_start
|
|
|
|
return best_command_label, best_command_confidence, best_overall_probs, best_command_start
|
|
|
|
|
|
def print_topk(probs: np.ndarray, topk: int) -> None:
|
|
top_indices = np.argsort(probs[: len(LABELS)])[::-1][:topk]
|
|
print("[INFO] Top classes:")
|
|
for index in top_indices:
|
|
print(f" {LABELS[int(index)]:>7s}: {float(probs[int(index)]):.6f}")
|
|
|
|
|
|
def parse_args() -> argparse.Namespace:
|
|
parser = argparse.ArgumentParser(
|
|
description="Run Tom game's tiny KWS TFLite model through Python."
|
|
)
|
|
parser.add_argument("--model", type=Path, default=DEFAULT_MODEL, help="Path to .tflite model")
|
|
parser.add_argument("--wav", type=Path, help="16-bit PCM WAV input")
|
|
parser.add_argument("--record", type=float, help="Record N seconds from microphone")
|
|
parser.add_argument("--device", help="sounddevice input device index/name for --record")
|
|
parser.add_argument("--mic-rate", type=int, default=48000, help="Microphone sample rate")
|
|
parser.add_argument("--channels", type=int, default=1, help="Input channel count")
|
|
parser.add_argument("--save-wav", type=Path, help="Save recorded raw microphone audio")
|
|
parser.add_argument("--threshold", type=float, default=0.75, help="Command confidence threshold")
|
|
parser.add_argument("--topk", type=int, default=5, help="Number of classes to print")
|
|
parser.add_argument("--no-trim", action="store_true", help="Disable silence trimming before KWS")
|
|
parser.add_argument("--list-devices", action="store_true", help="List sounddevice devices and exit")
|
|
return parser.parse_args()
|
|
|
|
|
|
def main() -> int:
|
|
args = parse_args()
|
|
|
|
try:
|
|
if args.list_devices:
|
|
list_devices()
|
|
return 0
|
|
|
|
if args.wav is None and args.record is None:
|
|
raise ValueError("Specify either --wav or --record, or use --list-devices")
|
|
|
|
if args.wav is not None and args.record is not None:
|
|
raise ValueError("Use only one input mode: --wav or --record")
|
|
|
|
model_path = args.model.resolve()
|
|
if not model_path.exists():
|
|
raise FileNotFoundError(f"Model not found: {model_path}")
|
|
|
|
backend, interpreter = load_interpreter(model_path)
|
|
print(f"[INFO] Interpreter backend: {backend}")
|
|
print(f"[INFO] Model: {model_path}")
|
|
print(f"[INFO] Input details: {interpreter.get_input_details()[0]}")
|
|
print(f"[INFO] Output details: {interpreter.get_output_details()[0]}")
|
|
|
|
if args.wav is not None:
|
|
samples, sample_rate, channels = read_wav_int16(args.wav.resolve())
|
|
print(f"[INFO] WAV: {args.wav.resolve()} ({sample_rate} Hz, {channels} channel(s))")
|
|
else:
|
|
sample_rate = args.mic_rate
|
|
channels = args.channels
|
|
samples = record_audio(args.record, sample_rate, channels, args.device)
|
|
if args.save_wav is not None:
|
|
write_wav_int16(args.save_wav.resolve(), samples, sample_rate, channels)
|
|
print(f"[INFO] Saved recording: {args.save_wav.resolve()}")
|
|
|
|
label, confidence, probs, start = recognize(
|
|
interpreter,
|
|
samples,
|
|
sample_rate,
|
|
channels,
|
|
args.threshold,
|
|
trim=not args.no_trim,
|
|
)
|
|
|
|
print_topk(probs, max(1, args.topk))
|
|
print(
|
|
f"[RESULT] command={label} confidence={confidence:.6f} "
|
|
f"window_start_sample={start} threshold={args.threshold:.3f}"
|
|
)
|
|
return 0
|
|
|
|
except Exception as exc:
|
|
print(f"[ERROR] {exc}", file=sys.stderr)
|
|
return 2
|
|
|
|
|
|
if __name__ == "__main__":
|
|
raise SystemExit(main())
|