Files
IMX6U-Game/src/Apps/Game/tools/kws_python_test.py
T

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())