完成语音识别功能(win测试,arm未知)
This commit is contained in:
Binary file not shown.
Binary file not shown.
@@ -40,6 +40,12 @@ namespace
|
||||
{ "Tom-say2.png", "tom_say2", ResizeMode::Fit, 700, 450 },
|
||||
{ "Tom-say3.png", "tom_say3", ResizeMode::Fit, 700, 450 },
|
||||
{ "Tom-say4.png", "tom_say4", ResizeMode::Fit, 700, 450 },
|
||||
{ "Tom-jump1.png", "tom_jump1", ResizeMode::Fit, 700, 450 },
|
||||
{ "Tom-jump2.png", "tom_jump2", ResizeMode::Fit, 700, 450 },
|
||||
{ "Tom-jump3.png", "tom_jump3", ResizeMode::Fit, 700, 450 },
|
||||
{ "Tom-jump4.png", "tom_jump4", ResizeMode::Fit, 700, 450 },
|
||||
{ "Tom-jump5.png", "tom_jump5", ResizeMode::Fit, 700, 450 },
|
||||
{ "Tom-jump6.png", "tom_jump6", ResizeMode::Fit, 700, 450 },
|
||||
{ "Tom-stand.png", "tom_stand", ResizeMode::Fit, 700, 450 },
|
||||
{ "ui-fat.png", "ui_fat", ResizeMode::Fit, 90, 90 },
|
||||
{ "ui-hand.png", "ui_hand", ResizeMode::Fit, 90, 90 },
|
||||
|
||||
@@ -0,0 +1,477 @@
|
||||
#!/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())
|
||||
Reference in New Issue
Block a user