#!/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=" 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=" 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())