Compare commits
9 Commits
ae52578ed7
..
main
| Author | SHA1 | Date | |
|---|---|---|---|
| 08b6543b79 | |||
| 6dee2a2ff3 | |||
| ea0c007441 | |||
| b45a8e2e85 | |||
| 1f6c3f3de8 | |||
| 37b85cd683 | |||
| c3f93e4441 | |||
| c612a7ad71 | |||
| e86c2301ec |
@@ -1,24 +0,0 @@
|
|||||||
# Server
|
|
||||||
POSEFIT_WS_HOST=0.0.0.0
|
|
||||||
POSEFIT_WS_PORT=8765
|
|
||||||
|
|
||||||
# Video processing
|
|
||||||
POSEFIT_PROCESS_EVERY_N_FRAMES=1
|
|
||||||
|
|
||||||
# Model
|
|
||||||
POSEFIT_MODEL_PATH=pose_models/pose_landmarker_full.task
|
|
||||||
POSEFIT_PREFER_GPU=1
|
|
||||||
|
|
||||||
# Dead bug exercise
|
|
||||||
POSEFIT_VISIBILITY_THRESHOLD=0.45
|
|
||||||
POSEFIT_EXTENSION_CONFIRM_FRAMES=4
|
|
||||||
POSEFIT_RESET_CONFIRM_FRAMES=3
|
|
||||||
|
|
||||||
# Audio
|
|
||||||
POSEFIT_REP_ANNOUNCER_ENABLED=1
|
|
||||||
POSEFIT_REP_ANNOUNCER_RATE=185
|
|
||||||
POSEFIT_REP_ANNOUNCER_VOLUME=1.0
|
|
||||||
|
|
||||||
# Logging
|
|
||||||
POSEFIT_LOG_ROTATION=20 MB
|
|
||||||
POSEFIT_LOG_RETENTION=14 days
|
|
||||||
@@ -2,3 +2,5 @@
|
|||||||
.idea/
|
.idea/
|
||||||
__pycache__/
|
__pycache__/
|
||||||
logs/
|
logs/
|
||||||
|
|
||||||
|
resources/
|
||||||
|
|||||||
@@ -7,8 +7,4 @@ Real-time exercise pose detection and coaching via WebRTC.
|
|||||||
```
|
```
|
||||||
pip install -r requirements.txt
|
pip install -r requirements.txt
|
||||||
python run.py
|
python run.py
|
||||||
```
|
``
|
||||||
|
|
||||||
## Configuration
|
|
||||||
|
|
||||||
Copy `.env.example` to `.env` and adjust settings, or set environment variables directly.
|
|
||||||
@@ -0,0 +1,227 @@
|
|||||||
|
# app/audio/generate.py
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import platform
|
||||||
|
import shutil
|
||||||
|
import subprocess
|
||||||
|
import wave
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from loguru import logger
|
||||||
|
|
||||||
|
|
||||||
|
def generate_rep_audio_files(
|
||||||
|
*,
|
||||||
|
max_count: int,
|
||||||
|
rate: int,
|
||||||
|
output_dir: Path,
|
||||||
|
overwrite: bool = False,
|
||||||
|
trim_leading_silence: bool = True,
|
||||||
|
trim_silence_threshold: int = 500,
|
||||||
|
trim_silence_padding_ms: int = 20,
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
确保 0~max_count 的运动次数语音 wav 文件存在。
|
||||||
|
|
||||||
|
默认生成到:
|
||||||
|
|
||||||
|
resources/audio/reps/0.aiff # macOS
|
||||||
|
resources/audio/reps/0.wav # Windows / Linux
|
||||||
|
...
|
||||||
|
resources/audio/reps/200.aiff 或 200.wav
|
||||||
|
|
||||||
|
服务启动时调用一次即可。
|
||||||
|
"""
|
||||||
|
output_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
system = platform.system().lower()
|
||||||
|
suffix = ".aiff" if system == "darwin" else ".wav"
|
||||||
|
|
||||||
|
missing_counts = [
|
||||||
|
count
|
||||||
|
for count in range(0, max_count + 1)
|
||||||
|
if overwrite or not _audio_path(output_dir, count, suffix=suffix).exists()
|
||||||
|
]
|
||||||
|
|
||||||
|
if not missing_counts:
|
||||||
|
logger.info("Rep audio files already prepared: {}", output_dir)
|
||||||
|
else:
|
||||||
|
logger.info(
|
||||||
|
"Preparing rep audio files, system={}, count={}, output_dir={}",
|
||||||
|
system,
|
||||||
|
len(missing_counts),
|
||||||
|
output_dir,
|
||||||
|
)
|
||||||
|
|
||||||
|
if system == "darwin":
|
||||||
|
_generate_with_macos_say(
|
||||||
|
counts=missing_counts,
|
||||||
|
output_dir=output_dir,
|
||||||
|
rate=rate,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
_generate_with_pyttsx3(
|
||||||
|
counts=missing_counts,
|
||||||
|
output_dir=output_dir,
|
||||||
|
rate=rate,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.info("Rep audio files prepared: {}", output_dir)
|
||||||
|
|
||||||
|
if trim_leading_silence and suffix == ".wav":
|
||||||
|
_trim_leading_silence_files(
|
||||||
|
counts=list(range(0, max_count + 1)),
|
||||||
|
output_dir=output_dir,
|
||||||
|
suffix=suffix,
|
||||||
|
threshold=trim_silence_threshold,
|
||||||
|
padding_ms=trim_silence_padding_ms,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _generate_with_macos_say(
|
||||||
|
*,
|
||||||
|
counts: list[int],
|
||||||
|
output_dir: Path,
|
||||||
|
rate: int,
|
||||||
|
) -> None:
|
||||||
|
"""macOS 使用 say 命令生成 wav。"""
|
||||||
|
if platform.system().lower() != "darwin":
|
||||||
|
raise RuntimeError("say command is only available on macOS")
|
||||||
|
|
||||||
|
if shutil.which("say") is None:
|
||||||
|
raise RuntimeError("macOS say command not found")
|
||||||
|
|
||||||
|
for count in counts:
|
||||||
|
audio_file = _audio_path(output_dir, count, suffix=".aiff")
|
||||||
|
|
||||||
|
try:
|
||||||
|
subprocess.run(
|
||||||
|
[
|
||||||
|
"say",
|
||||||
|
"-r",
|
||||||
|
str(rate),
|
||||||
|
"--file-format=AIFF",
|
||||||
|
"-o",
|
||||||
|
str(audio_file),
|
||||||
|
str(count),
|
||||||
|
],
|
||||||
|
stdout=subprocess.DEVNULL,
|
||||||
|
stderr=subprocess.PIPE,
|
||||||
|
text=True,
|
||||||
|
check=True,
|
||||||
|
)
|
||||||
|
except subprocess.CalledProcessError as exc:
|
||||||
|
message = exc.stderr.strip() or f"exit status {exc.returncode}"
|
||||||
|
raise RuntimeError(f"Failed to generate {audio_file}: {message}") from exc
|
||||||
|
|
||||||
|
|
||||||
|
def _generate_with_pyttsx3(
|
||||||
|
*,
|
||||||
|
counts: list[int],
|
||||||
|
output_dir: Path,
|
||||||
|
rate: int,
|
||||||
|
) -> None:
|
||||||
|
"""Windows / Linux 使用 pyttsx3 生成 wav。"""
|
||||||
|
try:
|
||||||
|
import pyttsx3
|
||||||
|
except Exception as exc:
|
||||||
|
raise RuntimeError(f"pyttsx3 unavailable: {exc}") from exc
|
||||||
|
|
||||||
|
engine = pyttsx3.init()
|
||||||
|
engine.setProperty("rate", rate)
|
||||||
|
engine.setProperty("volume", 1.0)
|
||||||
|
|
||||||
|
for count in counts:
|
||||||
|
audio_file = _audio_path(output_dir, count, suffix=".wav")
|
||||||
|
engine.save_to_file(str(count), str(audio_file))
|
||||||
|
|
||||||
|
engine.runAndWait()
|
||||||
|
|
||||||
|
|
||||||
|
def _audio_path(output_dir: Path, count: int, *, suffix: str) -> Path:
|
||||||
|
return output_dir / f"{count}{suffix}"
|
||||||
|
|
||||||
|
|
||||||
|
def _trim_leading_silence_files(
|
||||||
|
*,
|
||||||
|
counts: list[int],
|
||||||
|
output_dir: Path,
|
||||||
|
suffix: str,
|
||||||
|
threshold: int,
|
||||||
|
padding_ms: int,
|
||||||
|
) -> None:
|
||||||
|
trimmed = 0
|
||||||
|
total_removed_ms = 0.0
|
||||||
|
|
||||||
|
for count in counts:
|
||||||
|
audio_file = _audio_path(output_dir, count, suffix=suffix)
|
||||||
|
if not audio_file.exists():
|
||||||
|
continue
|
||||||
|
removed_ms = _trim_leading_silence(audio_file, threshold=threshold, padding_ms=padding_ms)
|
||||||
|
if removed_ms > 0:
|
||||||
|
trimmed += 1
|
||||||
|
total_removed_ms += removed_ms
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"Rep audio leading silence trim complete: files_trimmed={}, total_removed_ms={:.1f}, threshold={}, padding_ms={}",
|
||||||
|
trimmed,
|
||||||
|
total_removed_ms,
|
||||||
|
threshold,
|
||||||
|
padding_ms,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _trim_leading_silence(audio_file: Path, *, threshold: int, padding_ms: int) -> float:
|
||||||
|
with wave.open(str(audio_file), "rb") as reader:
|
||||||
|
params = reader.getparams()
|
||||||
|
frames = reader.readframes(params.nframes)
|
||||||
|
|
||||||
|
frame_size = params.sampwidth * params.nchannels
|
||||||
|
if params.nframes <= 0 or frame_size <= 0:
|
||||||
|
return 0.0
|
||||||
|
|
||||||
|
chunk_frames = max(1, params.framerate // 100)
|
||||||
|
leading_frames = 0
|
||||||
|
offset = 0
|
||||||
|
chunk_size = chunk_frames * frame_size
|
||||||
|
|
||||||
|
while offset < len(frames):
|
||||||
|
chunk = frames[offset : offset + chunk_size]
|
||||||
|
if _pcm_rms(chunk, params.sampwidth) > threshold:
|
||||||
|
break
|
||||||
|
chunk_frame_count = len(chunk) // frame_size
|
||||||
|
leading_frames += chunk_frame_count
|
||||||
|
offset += chunk_size
|
||||||
|
|
||||||
|
padding_frames = int(params.framerate * max(0, padding_ms) / 1000)
|
||||||
|
remove_frames = max(0, leading_frames - padding_frames)
|
||||||
|
if remove_frames <= 0:
|
||||||
|
return 0.0
|
||||||
|
|
||||||
|
start = min(len(frames), remove_frames * frame_size)
|
||||||
|
trimmed_frames = frames[start:]
|
||||||
|
if not trimmed_frames:
|
||||||
|
return 0.0
|
||||||
|
|
||||||
|
with wave.open(str(audio_file), "wb") as writer:
|
||||||
|
writer.setparams(params)
|
||||||
|
writer.writeframes(trimmed_frames)
|
||||||
|
|
||||||
|
return remove_frames / params.framerate * 1000
|
||||||
|
|
||||||
|
|
||||||
|
def _pcm_rms(chunk: bytes, sample_width: int) -> float:
|
||||||
|
if not chunk:
|
||||||
|
return 0.0
|
||||||
|
|
||||||
|
if sample_width == 2:
|
||||||
|
sample_count = len(chunk) // 2
|
||||||
|
if sample_count == 0:
|
||||||
|
return 0.0
|
||||||
|
total = 0
|
||||||
|
for i in range(0, sample_count * 2, 2):
|
||||||
|
sample = int.from_bytes(chunk[i : i + 2], "little", signed=True)
|
||||||
|
total += sample * sample
|
||||||
|
return (total / sample_count) ** 0.5
|
||||||
|
|
||||||
|
peak = max(abs(byte - 128) for byte in chunk)
|
||||||
|
return float(peak)
|
||||||
+138
-44
@@ -1,84 +1,178 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import queue
|
import queue
|
||||||
|
import shutil
|
||||||
import subprocess
|
import subprocess
|
||||||
import sys
|
import sys
|
||||||
import threading
|
import threading
|
||||||
from typing import Any
|
import time
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
|
|
||||||
class RepAnnouncer:
|
class RepAnnouncer:
|
||||||
def __init__(self, *, enabled: bool = True, rate: int = 185, volume: float = 1.0) -> None:
|
"""运动次数语音播报器:读取预生成的音频文件直接播放"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
enabled: bool = True,
|
||||||
|
max_count: int = 200,
|
||||||
|
audio_dir: str | Path = "resources/audio/reps",
|
||||||
|
) -> None:
|
||||||
self.enabled = enabled
|
self.enabled = enabled
|
||||||
self.rate = rate
|
self.max_count = max_count
|
||||||
self.volume = volume
|
self.audio_dir = Path(audio_dir)
|
||||||
self._queue: queue.Queue[str | None] = queue.Queue()
|
|
||||||
|
self._queue: queue.Queue[tuple[int, float] | None] = queue.Queue()
|
||||||
self._thread: threading.Thread | None = None
|
self._thread: threading.Thread | None = None
|
||||||
self._engine: Any | None = None
|
|
||||||
self._use_macos_say = False
|
|
||||||
self._current_process: subprocess.Popen | None = None
|
self._current_process: subprocess.Popen | None = None
|
||||||
|
self._closed = False
|
||||||
|
self._play_lock = threading.Lock()
|
||||||
|
|
||||||
|
self._platform = sys.platform
|
||||||
|
self._direct_playback = self._platform.startswith("win")
|
||||||
|
|
||||||
if self.enabled:
|
if self.enabled:
|
||||||
self._start()
|
self._start()
|
||||||
|
|
||||||
def announce_count(self, count: int) -> None:
|
def announce_count(self, count: int) -> None:
|
||||||
if not self.enabled or count <= 0:
|
"""将次数放入队列,后台线程播放对应音频"""
|
||||||
|
if not self.enabled or self._closed:
|
||||||
return
|
return
|
||||||
while True:
|
if count <= 0 or count > self.max_count:
|
||||||
|
return
|
||||||
|
|
||||||
|
requested_at = time.perf_counter()
|
||||||
|
if self._direct_playback:
|
||||||
|
audio_file = self._audio_path(count)
|
||||||
|
if not audio_file.exists():
|
||||||
|
logger.warning("Rep audio file missing: {}", audio_file)
|
||||||
|
return
|
||||||
try:
|
try:
|
||||||
self._queue.get_nowait()
|
self._play(audio_file)
|
||||||
except queue.Empty:
|
logger.info(
|
||||||
break
|
"Rep audio submitted immediately: count={}, submit_ms={:.1f}",
|
||||||
self._queue.put(str(count))
|
count,
|
||||||
|
(time.perf_counter() - requested_at) * 1000,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("Failed to play rep count {}: {}", count, exc)
|
||||||
|
return
|
||||||
|
|
||||||
|
self._clear_pending_counts()
|
||||||
|
self._queue.put((count, requested_at))
|
||||||
|
logger.info("Rep audio queued: count={}", count)
|
||||||
|
|
||||||
def close(self) -> None:
|
def close(self) -> None:
|
||||||
if not self.enabled:
|
"""停止播报线程并释放资源"""
|
||||||
|
if not self.enabled or self._closed:
|
||||||
return
|
return
|
||||||
|
|
||||||
|
self._closed = True
|
||||||
self._queue.put(None)
|
self._queue.put(None)
|
||||||
|
|
||||||
if self._thread is not None:
|
if self._thread is not None:
|
||||||
self._thread.join(timeout=1.0)
|
self._thread.join(timeout=1.0)
|
||||||
if self._current_process is not None and self._current_process.poll() is None:
|
|
||||||
self._current_process.terminate()
|
self._stop_current_playback()
|
||||||
|
logger.info("Rep announcer closed")
|
||||||
|
|
||||||
def _start(self) -> None:
|
def _start(self) -> None:
|
||||||
if sys.platform == "darwin":
|
"""启动后台播报线程"""
|
||||||
self._use_macos_say = True
|
self.audio_dir.mkdir(parents=True, exist_ok=True)
|
||||||
logger.info("Rep announcer initialized with macOS say")
|
|
||||||
else:
|
|
||||||
try:
|
|
||||||
import pyttsx3
|
|
||||||
|
|
||||||
self._engine = pyttsx3.init()
|
if self._direct_playback:
|
||||||
self._engine.setProperty("rate", self.rate)
|
import winsound
|
||||||
self._engine.setProperty("volume", self.volume)
|
|
||||||
logger.info("Rep announcer initialized with pyttsx3")
|
logger.info("Rep announcer initialized in direct Windows mode, audio_dir={}", self.audio_dir)
|
||||||
except Exception as exc:
|
return
|
||||||
self.enabled = False
|
|
||||||
logger.warning("Rep announcer disabled, pyttsx3 unavailable: {}", exc)
|
|
||||||
return
|
|
||||||
|
|
||||||
self._thread = threading.Thread(target=self._run, name="RepAnnouncer", daemon=True)
|
self._thread = threading.Thread(target=self._run, name="RepAnnouncer", daemon=True)
|
||||||
self._thread.start()
|
self._thread.start()
|
||||||
|
|
||||||
|
logger.info("Rep announcer initialized in queued mode, audio_dir={}", self.audio_dir)
|
||||||
|
|
||||||
def _run(self) -> None:
|
def _run(self) -> None:
|
||||||
|
"""后台线程:从队列取次数,播放对应音频文件"""
|
||||||
while True:
|
while True:
|
||||||
text = self._queue.get()
|
item = self._queue.get()
|
||||||
if text is None:
|
if item is None:
|
||||||
return
|
return
|
||||||
|
count, requested_at = item
|
||||||
|
|
||||||
|
audio_file = self._audio_path(count)
|
||||||
|
if not audio_file.exists():
|
||||||
|
logger.warning("Rep audio file missing: {}", audio_file)
|
||||||
|
continue
|
||||||
|
|
||||||
try:
|
try:
|
||||||
if self._use_macos_say:
|
self._play(audio_file)
|
||||||
if self._current_process is not None and self._current_process.poll() is None:
|
logger.info(
|
||||||
self._current_process.terminate()
|
"Rep audio submitted from queue: count={}, queue_ms={:.1f}",
|
||||||
self._current_process = subprocess.Popen(
|
count,
|
||||||
["say", "-r", str(self.rate), text],
|
(time.perf_counter() - requested_at) * 1000,
|
||||||
stdout=subprocess.DEVNULL,
|
)
|
||||||
stderr=subprocess.DEVNULL,
|
|
||||||
)
|
|
||||||
elif self._engine is not None:
|
|
||||||
self._engine.say(text)
|
|
||||||
self._engine.runAndWait()
|
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.warning("Failed to announce rep count {}: {}", text, exc)
|
logger.warning("Failed to play rep count {}: {}", count, exc)
|
||||||
|
|
||||||
|
def _play(self, audio_file: Path) -> None:
|
||||||
|
"""播放音频文件(平台自适应)"""
|
||||||
|
with self._play_lock:
|
||||||
|
self._stop_current_playback()
|
||||||
|
|
||||||
|
if self._platform == "darwin":
|
||||||
|
self._current_process = subprocess.Popen(
|
||||||
|
["afplay", str(audio_file)],
|
||||||
|
stdout=subprocess.DEVNULL,
|
||||||
|
stderr=subprocess.DEVNULL,
|
||||||
|
)
|
||||||
|
elif self._platform.startswith("win"):
|
||||||
|
import winsound
|
||||||
|
|
||||||
|
winsound.PlaySound(
|
||||||
|
str(audio_file),
|
||||||
|
winsound.SND_FILENAME | winsound.SND_ASYNC | winsound.SND_NODEFAULT,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
player = shutil.which("paplay") or shutil.which("aplay")
|
||||||
|
if player is None:
|
||||||
|
logger.warning("No audio player found")
|
||||||
|
return
|
||||||
|
self._current_process = subprocess.Popen(
|
||||||
|
[player, str(audio_file)],
|
||||||
|
stdout=subprocess.DEVNULL,
|
||||||
|
stderr=subprocess.DEVNULL,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _stop_current_playback(self) -> None:
|
||||||
|
"""中断当前正在播放的声音"""
|
||||||
|
if self._platform.startswith("win"):
|
||||||
|
try:
|
||||||
|
import winsound
|
||||||
|
|
||||||
|
winsound.PlaySound(None, winsound.SND_PURGE)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return
|
||||||
|
|
||||||
|
if self._current_process is not None and self._current_process.poll() is None:
|
||||||
|
self._current_process.terminate()
|
||||||
|
self._current_process = None
|
||||||
|
|
||||||
|
def _audio_path(self, count: int) -> Path:
|
||||||
|
"""获取某个次数对应的音频文件路径"""
|
||||||
|
suffix = ".aiff" if self._platform == "darwin" else ".wav"
|
||||||
|
return self.audio_dir / f"{count}{suffix}"
|
||||||
|
|
||||||
|
def _clear_pending_counts(self) -> None:
|
||||||
|
"""清空队列中等待播放的次数,避免语音堆积"""
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
item = self._queue.get_nowait()
|
||||||
|
if item is None:
|
||||||
|
self._queue.put(None)
|
||||||
|
return
|
||||||
|
except queue.Empty:
|
||||||
|
return
|
||||||
|
|||||||
+13
-1
@@ -2,9 +2,21 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from app.diagnostics.crash_handler import enable_crash_handler
|
from app.diagnostics.crash_handler import enable_crash_handler
|
||||||
from configs.load import config
|
from configs.load import config
|
||||||
|
from app.audio.generate import generate_rep_audio_files
|
||||||
|
|
||||||
def startup() -> None:
|
def startup() -> None:
|
||||||
|
"""应用启动初始化:开启崩溃日志和日志系统"""
|
||||||
enable_crash_handler(config.logging.dir_path)
|
enable_crash_handler(config.logging.dir_path)
|
||||||
from app.core.logging import setup_logging
|
from app.core.logging import setup_logging
|
||||||
setup_logging()
|
setup_logging()
|
||||||
|
|
||||||
|
# 生成运动次数语音文件
|
||||||
|
generate_rep_audio_files(
|
||||||
|
max_count=config.audio.rep_max_count,
|
||||||
|
rate=config.audio.rep_announcer_rate,
|
||||||
|
output_dir=config.audio.resolved_audio_dir,
|
||||||
|
overwrite=False,
|
||||||
|
trim_leading_silence=config.audio.trim_leading_silence,
|
||||||
|
trim_silence_threshold=config.audio.trim_silence_threshold,
|
||||||
|
trim_silence_padding_ms=config.audio.trim_silence_padding_ms,
|
||||||
|
)
|
||||||
|
|||||||
+2
-1
@@ -8,11 +8,12 @@ from configs.load import config
|
|||||||
|
|
||||||
|
|
||||||
def setup_logging() -> None:
|
def setup_logging() -> None:
|
||||||
|
"""配置loguru日志输出到按日期轮转的日志文件"""
|
||||||
log_dir = config.logging.dir_path
|
log_dir = config.logging.dir_path
|
||||||
log_dir.mkdir(parents=True, exist_ok=True)
|
log_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
logger.add(
|
logger.add(
|
||||||
log_dir / "posefit-server_{time:YYYY-MM-DD}.log",
|
log_dir /"{time:YYYY-MM-DD}.log",
|
||||||
rotation=config.logging.rotation,
|
rotation=config.logging.rotation,
|
||||||
retention=config.logging.retention,
|
retention=config.logging.retention,
|
||||||
enqueue=True,
|
enqueue=True,
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ from pathlib import Path
|
|||||||
|
|
||||||
|
|
||||||
def enable_crash_handler(log_dir: str | Path) -> None:
|
def enable_crash_handler(log_dir: str | Path) -> None:
|
||||||
|
"""启用faulthandler,将崩溃堆栈写入日志文件"""
|
||||||
log_dir = Path(log_dir)
|
log_dir = Path(log_dir)
|
||||||
log_dir.mkdir(parents=True, exist_ok=True)
|
log_dir.mkdir(parents=True, exist_ok=True)
|
||||||
crash_log = open(log_dir / "posefit-crash.log", "a", buffering=1)
|
crash_log = open(log_dir / "posefit-crash.log", "a", buffering=1)
|
||||||
|
|||||||
@@ -5,28 +5,33 @@ from contextlib import contextmanager
|
|||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
|
|
||||||
class PerfTimer:
|
class PerfTimer:
|
||||||
|
"""性能计时器,用于测量代码段执行耗时"""
|
||||||
|
|
||||||
def __init__(self, name: str = "") -> None:
|
def __init__(self, name: str = "") -> None:
|
||||||
self.name = name
|
self.name = name
|
||||||
self._start = 0.0
|
self._start = 0.0
|
||||||
self._elapsed = 0.0
|
self._elapsed = 0.0
|
||||||
|
|
||||||
def start(self) -> PerfTimer:
|
def start(self) -> PerfTimer:
|
||||||
|
"""启动计时器"""
|
||||||
self._start = time.perf_counter()
|
self._start = time.perf_counter()
|
||||||
return self
|
return self
|
||||||
|
|
||||||
def stop(self) -> float:
|
def stop(self) -> float:
|
||||||
|
"""停止计时器并返回耗时(秒)"""
|
||||||
self._elapsed = time.perf_counter() - self._start
|
self._elapsed = time.perf_counter() - self._start
|
||||||
return self._elapsed
|
return self._elapsed
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def elapsed_ms(self) -> float:
|
def elapsed_ms(self) -> float:
|
||||||
|
"""返回已记录耗时(毫秒)"""
|
||||||
return self._elapsed * 1000
|
return self._elapsed * 1000
|
||||||
|
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def measure(name: str = ""):
|
def measure(name: str = ""):
|
||||||
|
"""上下文管理器:进入时计时,退出时记录耗时日志"""
|
||||||
timer = PerfTimer(name).start()
|
timer = PerfTimer(name).start()
|
||||||
yield timer
|
yield timer
|
||||||
elapsed = timer.stop()
|
elapsed = timer.stop()
|
||||||
|
|||||||
@@ -32,8 +32,9 @@ from app.vision.pose_types import (
|
|||||||
RIGHT_WRIST,
|
RIGHT_WRIST,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
class DeadBugDetector:
|
class DeadBugDetector:
|
||||||
|
"""死虫式(Dead Bug)运动检测器"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -43,6 +44,7 @@ class DeadBugDetector:
|
|||||||
reset_confirm_frames: int = 3,
|
reset_confirm_frames: int = 3,
|
||||||
prefer_gpu: bool = True,
|
prefer_gpu: bool = True,
|
||||||
) -> None:
|
) -> None:
|
||||||
|
"""初始化姿态检测器、状态机和可视化渲染组件"""
|
||||||
self.visibility_threshold = visibility_threshold
|
self.visibility_threshold = visibility_threshold
|
||||||
|
|
||||||
self._latest_result = None
|
self._latest_result = None
|
||||||
@@ -50,6 +52,7 @@ class DeadBugDetector:
|
|||||||
self._result_event = threading.Event()
|
self._result_event = threading.Event()
|
||||||
self._inflight = False
|
self._inflight = False
|
||||||
self._inflight_started_at = 0.0
|
self._inflight_started_at = 0.0
|
||||||
|
self.last_timing: dict[str, float | bool] = {}
|
||||||
|
|
||||||
def on_result(pose_result, _image, _timestamp_ms):
|
def on_result(pose_result, _image, _timestamp_ms):
|
||||||
with self._result_lock:
|
with self._result_lock:
|
||||||
@@ -72,10 +75,14 @@ class DeadBugDetector:
|
|||||||
self._last_timestamp_ms = -1
|
self._last_timestamp_ms = -1
|
||||||
|
|
||||||
def close(self) -> None:
|
def close(self) -> None:
|
||||||
|
"""释放MediaPipe模型资源"""
|
||||||
self._landmarker.close()
|
self._landmarker.close()
|
||||||
|
|
||||||
def process_frame(self, bgr_frame: np.ndarray, timestamp_ms: int) -> tuple[np.ndarray, DeadBugResult]:
|
def process_frame(self, bgr_frame: np.ndarray, timestamp_ms: int) -> tuple[np.ndarray, DeadBugResult]:
|
||||||
|
"""处理单帧:姿态检测、指标计算、状态机更新、可视化叠加"""
|
||||||
|
total_started = time.perf_counter()
|
||||||
timestamp_ms = self._normalize_timestamp(timestamp_ms)
|
timestamp_ms = self._normalize_timestamp(timestamp_ms)
|
||||||
|
normalize_done = time.perf_counter()
|
||||||
|
|
||||||
with self._result_lock:
|
with self._result_lock:
|
||||||
if self._inflight and time.monotonic() - self._inflight_started_at > 0.5:
|
if self._inflight and time.monotonic() - self._inflight_started_at > 0.5:
|
||||||
@@ -86,9 +93,11 @@ class DeadBugDetector:
|
|||||||
if should_submit:
|
if should_submit:
|
||||||
self._inflight = True
|
self._inflight = True
|
||||||
self._inflight_started_at = time.monotonic()
|
self._inflight_started_at = time.monotonic()
|
||||||
|
lock_done = time.perf_counter()
|
||||||
|
|
||||||
if should_submit:
|
if should_submit:
|
||||||
rgba_frame = bgr_to_rgba(bgr_frame)
|
rgba_frame = bgr_to_rgba(bgr_frame)
|
||||||
|
convert_done = time.perf_counter()
|
||||||
mp_image = mp.Image(image_format=mp.ImageFormat.SRGBA, data=rgba_frame)
|
mp_image = mp.Image(image_format=mp.ImageFormat.SRGBA, data=rgba_frame)
|
||||||
self._result_event.clear()
|
self._result_event.clear()
|
||||||
try:
|
try:
|
||||||
@@ -98,14 +107,23 @@ class DeadBugDetector:
|
|||||||
self._inflight = False
|
self._inflight = False
|
||||||
self._inflight_started_at = 0.0
|
self._inflight_started_at = 0.0
|
||||||
raise
|
raise
|
||||||
|
submit_done = time.perf_counter()
|
||||||
self._result_event.wait(timeout=0.08)
|
self._result_event.wait(timeout=0.08)
|
||||||
|
wait_done = time.perf_counter()
|
||||||
|
else:
|
||||||
|
convert_done = lock_done
|
||||||
|
submit_done = lock_done
|
||||||
|
wait_done = lock_done
|
||||||
|
|
||||||
with self._result_lock:
|
with self._result_lock:
|
||||||
pose_result = self._latest_result
|
pose_result = self._latest_result
|
||||||
|
result_read_done = time.perf_counter()
|
||||||
|
|
||||||
annotated = bgr_frame.copy()
|
annotated = bgr_frame.copy()
|
||||||
|
copy_done = time.perf_counter()
|
||||||
|
|
||||||
if pose_result is None or not pose_result.pose_landmarks:
|
if pose_result is None or not pose_result.pose_landmarks:
|
||||||
|
self._state.mark_no_pose()
|
||||||
result = DeadBugResult(
|
result = DeadBugResult(
|
||||||
rep_count=self._state.rep_count,
|
rep_count=self._state.rep_count,
|
||||||
phase=DeadBugPhase.NO_POSE,
|
phase=DeadBugPhase.NO_POSE,
|
||||||
@@ -115,12 +133,25 @@ class DeadBugDetector:
|
|||||||
metrics=None,
|
metrics=None,
|
||||||
)
|
)
|
||||||
draw_status_overlay(annotated, result)
|
draw_status_overlay(annotated, result)
|
||||||
|
self._record_timing(
|
||||||
|
total_started,
|
||||||
|
normalize_done,
|
||||||
|
lock_done,
|
||||||
|
convert_done,
|
||||||
|
submit_done,
|
||||||
|
wait_done,
|
||||||
|
result_read_done,
|
||||||
|
copy_done,
|
||||||
|
time.perf_counter(),
|
||||||
|
should_submit,
|
||||||
|
)
|
||||||
return annotated, result
|
return annotated, result
|
||||||
|
|
||||||
landmarks = [Point(lm.x, lm.y, lm.z, getattr(lm, "visibility", 1.0)) for lm in pose_result.pose_landmarks[0]]
|
landmarks = [Point(lm.x, lm.y, lm.z, getattr(lm, "visibility", 1.0)) for lm in pose_result.pose_landmarks[0]]
|
||||||
draw_landmarks(annotated, landmarks, REQUIRED_LANDMARKS, visibility_threshold=self.visibility_threshold)
|
draw_landmarks(annotated, landmarks, REQUIRED_LANDMARKS, visibility_threshold=self.visibility_threshold)
|
||||||
|
|
||||||
if not has_required_visibility(landmarks, REQUIRED_LANDMARKS, self.visibility_threshold):
|
if not has_required_visibility(landmarks, REQUIRED_LANDMARKS, self.visibility_threshold):
|
||||||
|
self._state.mark_no_pose()
|
||||||
result = DeadBugResult(
|
result = DeadBugResult(
|
||||||
rep_count=self._state.rep_count,
|
rep_count=self._state.rep_count,
|
||||||
phase=DeadBugPhase.NO_POSE,
|
phase=DeadBugPhase.NO_POSE,
|
||||||
@@ -130,6 +161,18 @@ class DeadBugDetector:
|
|||||||
metrics=None,
|
metrics=None,
|
||||||
)
|
)
|
||||||
draw_status_overlay(annotated, result)
|
draw_status_overlay(annotated, result)
|
||||||
|
self._record_timing(
|
||||||
|
total_started,
|
||||||
|
normalize_done,
|
||||||
|
lock_done,
|
||||||
|
convert_done,
|
||||||
|
submit_done,
|
||||||
|
wait_done,
|
||||||
|
result_read_done,
|
||||||
|
copy_done,
|
||||||
|
time.perf_counter(),
|
||||||
|
should_submit,
|
||||||
|
)
|
||||||
return annotated, result
|
return annotated, result
|
||||||
|
|
||||||
raw = calculate_metrics(
|
raw = calculate_metrics(
|
||||||
@@ -163,9 +206,48 @@ class DeadBugDetector:
|
|||||||
|
|
||||||
result = self._state.update(metrics)
|
result = self._state.update(metrics)
|
||||||
draw_status_overlay(annotated, result)
|
draw_status_overlay(annotated, result)
|
||||||
|
self._record_timing(
|
||||||
|
total_started,
|
||||||
|
normalize_done,
|
||||||
|
lock_done,
|
||||||
|
convert_done,
|
||||||
|
submit_done,
|
||||||
|
wait_done,
|
||||||
|
result_read_done,
|
||||||
|
copy_done,
|
||||||
|
time.perf_counter(),
|
||||||
|
should_submit,
|
||||||
|
)
|
||||||
return annotated, result
|
return annotated, result
|
||||||
|
|
||||||
|
def _record_timing(
|
||||||
|
self,
|
||||||
|
total_started: float,
|
||||||
|
normalize_done: float,
|
||||||
|
lock_done: float,
|
||||||
|
convert_done: float,
|
||||||
|
submit_done: float,
|
||||||
|
wait_done: float,
|
||||||
|
result_read_done: float,
|
||||||
|
copy_done: float,
|
||||||
|
finished: float,
|
||||||
|
submitted: bool,
|
||||||
|
) -> None:
|
||||||
|
self.last_timing = {
|
||||||
|
"total_ms": (finished - total_started) * 1000,
|
||||||
|
"timestamp_ms": (normalize_done - total_started) * 1000,
|
||||||
|
"lock_ms": (lock_done - normalize_done) * 1000,
|
||||||
|
"convert_ms": (convert_done - lock_done) * 1000,
|
||||||
|
"submit_ms": (submit_done - convert_done) * 1000,
|
||||||
|
"wait_ms": (wait_done - submit_done) * 1000,
|
||||||
|
"result_read_ms": (result_read_done - wait_done) * 1000,
|
||||||
|
"copy_ms": (copy_done - result_read_done) * 1000,
|
||||||
|
"postprocess_draw_ms": (finished - copy_done) * 1000,
|
||||||
|
"submitted": submitted,
|
||||||
|
}
|
||||||
|
|
||||||
def _normalize_timestamp(self, timestamp_ms: int) -> int:
|
def _normalize_timestamp(self, timestamp_ms: int) -> int:
|
||||||
|
"""确保时间戳严格递增(MediaPipe要求)"""
|
||||||
if timestamp_ms <= self._last_timestamp_ms:
|
if timestamp_ms <= self._last_timestamp_ms:
|
||||||
timestamp_ms = self._last_timestamp_ms + 1
|
timestamp_ms = self._last_timestamp_ms + 1
|
||||||
self._last_timestamp_ms = timestamp_ms
|
self._last_timestamp_ms = timestamp_ms
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ from app.exercises.dead_bug.types import Point
|
|||||||
|
|
||||||
|
|
||||||
def angle(a: Point, b: Point, c: Point) -> float:
|
def angle(a: Point, b: Point, c: Point) -> float:
|
||||||
|
"""计算以b为顶点的三点夹角(度数)"""
|
||||||
ba = np.array([a.x - b.x, a.y - b.y], dtype=np.float32)
|
ba = np.array([a.x - b.x, a.y - b.y], dtype=np.float32)
|
||||||
bc = np.array([c.x - b.x, c.y - b.y], dtype=np.float32)
|
bc = np.array([c.x - b.x, c.y - b.y], dtype=np.float32)
|
||||||
denom = float(np.linalg.norm(ba) * np.linalg.norm(bc))
|
denom = float(np.linalg.norm(ba) * np.linalg.norm(bc))
|
||||||
@@ -17,6 +18,7 @@ def angle(a: Point, b: Point, c: Point) -> float:
|
|||||||
|
|
||||||
|
|
||||||
def distance(a: Point, b: Point) -> float:
|
def distance(a: Point, b: Point) -> float:
|
||||||
|
"""计算两点之间的欧几里得距离(归一化坐标空间)"""
|
||||||
return float(np.hypot(a.x - b.x, a.y - b.y))
|
return float(np.hypot(a.x - b.x, a.y - b.y))
|
||||||
|
|
||||||
|
|
||||||
@@ -37,6 +39,7 @@ def calculate_metrics(
|
|||||||
right_ankle: int,
|
right_ankle: int,
|
||||||
visibility_threshold: float = 0.45,
|
visibility_threshold: float = 0.45,
|
||||||
) -> dict:
|
) -> dict:
|
||||||
|
"""计算四肢关节角度、伸展状态及反馈信息"""
|
||||||
left_elbow_angle = angle(lm[left_shoulder], lm[left_elbow], lm[left_wrist])
|
left_elbow_angle = angle(lm[left_shoulder], lm[left_elbow], lm[left_wrist])
|
||||||
right_elbow_angle = angle(lm[right_shoulder], lm[right_elbow], lm[right_wrist])
|
right_elbow_angle = angle(lm[right_shoulder], lm[right_elbow], lm[right_wrist])
|
||||||
left_knee_angle = angle(lm[left_hip], lm[left_knee], lm[left_ankle])
|
left_knee_angle = angle(lm[left_hip], lm[left_knee], lm[left_ankle])
|
||||||
|
|||||||
@@ -4,10 +4,12 @@ from app.exercises.dead_bug.types import DeadBugMetrics, Point
|
|||||||
|
|
||||||
|
|
||||||
def has_required_visibility(landmarks: list[Point], required_indices: tuple[int, ...], visibility_threshold: float) -> bool:
|
def has_required_visibility(landmarks: list[Point], required_indices: tuple[int, ...], visibility_threshold: float) -> bool:
|
||||||
|
"""检查所有必需关键点的可见度是否高于阈值"""
|
||||||
return all(landmarks[i].visibility >= visibility_threshold for i in required_indices)
|
return all(landmarks[i].visibility >= visibility_threshold for i in required_indices)
|
||||||
|
|
||||||
|
|
||||||
def detect_diagonal_extension(metrics: DeadBugMetrics) -> str | None:
|
def detect_diagonal_extension(metrics: DeadBugMetrics) -> str | None:
|
||||||
|
"""检测对角伸展(腿部只允许单侧伸展,手臂允许准备位上举带来的识别重叠)"""
|
||||||
if metrics.left_leg_extended and metrics.right_leg_extended:
|
if metrics.left_leg_extended and metrics.right_leg_extended:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -19,6 +21,7 @@ def detect_diagonal_extension(metrics: DeadBugMetrics) -> str | None:
|
|||||||
|
|
||||||
|
|
||||||
def is_ready_position(metrics: DeadBugMetrics) -> bool:
|
def is_ready_position(metrics: DeadBugMetrics) -> bool:
|
||||||
|
"""判断是否处于准备姿态(双膝弯曲且双腿未伸展;dead bug 准备位允许手臂上举)"""
|
||||||
knees_bent = metrics.left_knee_angle <= 140 and metrics.right_knee_angle <= 140
|
knees_bent = metrics.left_knee_angle <= 140 and metrics.right_knee_angle <= 140
|
||||||
legs_not_extended = not metrics.left_leg_extended and not metrics.right_leg_extended
|
legs_not_extended = not metrics.left_leg_extended and not metrics.right_leg_extended
|
||||||
return knees_bent and legs_not_extended and detect_diagonal_extension(metrics) is None
|
return knees_bent and legs_not_extended
|
||||||
|
|||||||
@@ -3,9 +3,18 @@ from __future__ import annotations
|
|||||||
from app.exercises.dead_bug.rules import detect_diagonal_extension, is_ready_position
|
from app.exercises.dead_bug.rules import detect_diagonal_extension, is_ready_position
|
||||||
from app.exercises.dead_bug.types import DeadBugMetrics, DeadBugPhase, DeadBugResult
|
from app.exercises.dead_bug.types import DeadBugMetrics, DeadBugPhase, DeadBugResult
|
||||||
|
|
||||||
|
_SMOOTHING_ALPHA = 0.45
|
||||||
|
_TREND_DELTA_DEGREES = 1.5
|
||||||
|
_SIDE_ANGLE_MARGIN = 10.0
|
||||||
|
_EXTENSION_START_ANGLE = 125.0
|
||||||
|
_EXTENSION_PEAK_ANGLE = 150.0
|
||||||
|
_READY_KNEE_ANGLE = 140.0
|
||||||
|
|
||||||
class DeadBugStateMachine:
|
class DeadBugStateMachine:
|
||||||
|
"""死虫式动作状态机:管理READY/EXTENDING/NEED_RESET/NO_POSE状态转换"""
|
||||||
|
|
||||||
def __init__(self, *, extension_confirm_frames: int = 4, reset_confirm_frames: int = 3) -> None:
|
def __init__(self, *, extension_confirm_frames: int = 4, reset_confirm_frames: int = 3) -> None:
|
||||||
|
"""初始化并设置状态转换确认帧数"""
|
||||||
self.extension_confirm_frames = extension_confirm_frames
|
self.extension_confirm_frames = extension_confirm_frames
|
||||||
self.reset_confirm_frames = reset_confirm_frames
|
self.reset_confirm_frames = reset_confirm_frames
|
||||||
|
|
||||||
@@ -15,10 +24,26 @@ class DeadBugStateMachine:
|
|||||||
self._candidate_side: str | None = None
|
self._candidate_side: str | None = None
|
||||||
self._candidate_frames = 0
|
self._candidate_frames = 0
|
||||||
self._reset_frames = 0
|
self._reset_frames = 0
|
||||||
|
self._smooth_left_knee_angle: float | None = None
|
||||||
|
self._smooth_right_knee_angle: float | None = None
|
||||||
|
self._left_knee_delta = 0.0
|
||||||
|
self._right_knee_delta = 0.0
|
||||||
|
|
||||||
|
def mark_no_pose(self) -> None:
|
||||||
|
"""姿态丢失时清掉候选帧;已确认的半程动作保留,等待重新可见后完成回收。"""
|
||||||
|
if self.phase == DeadBugPhase.READY:
|
||||||
|
self.phase = DeadBugPhase.NO_POSE
|
||||||
|
self.active_side = None
|
||||||
|
self._candidate_side = None
|
||||||
|
self._candidate_frames = 0
|
||||||
|
self._reset_frames = 0
|
||||||
|
|
||||||
def update(self, metrics: DeadBugMetrics) -> DeadBugResult:
|
def update(self, metrics: DeadBugMetrics) -> DeadBugResult:
|
||||||
side = detect_diagonal_extension(metrics)
|
"""根据传入指标更新状态机并返回本次结果"""
|
||||||
ready = is_ready_position(metrics)
|
self._update_knee_trends(metrics)
|
||||||
|
|
||||||
|
side = self._detect_motion_side(metrics)
|
||||||
|
ready = self._is_stable_ready(metrics)
|
||||||
|
|
||||||
if side is None:
|
if side is None:
|
||||||
self._candidate_side = None
|
self._candidate_side = None
|
||||||
@@ -30,13 +55,16 @@ class DeadBugStateMachine:
|
|||||||
self._candidate_frames = 1
|
self._candidate_frames = 1
|
||||||
|
|
||||||
if self.phase in (DeadBugPhase.READY, DeadBugPhase.NO_POSE):
|
if self.phase in (DeadBugPhase.READY, DeadBugPhase.NO_POSE):
|
||||||
|
if ready:
|
||||||
|
self.phase = DeadBugPhase.READY
|
||||||
if self._candidate_frames >= self.extension_confirm_frames and side is not None:
|
if self._candidate_frames >= self.extension_confirm_frames and side is not None:
|
||||||
self.phase = DeadBugPhase.EXTENDING
|
self.phase = DeadBugPhase.EXTENDING
|
||||||
self.active_side = side
|
self.active_side = side
|
||||||
self._reset_frames = 0
|
self._reset_frames = 0
|
||||||
elif self.phase == DeadBugPhase.EXTENDING:
|
elif self.phase == DeadBugPhase.EXTENDING:
|
||||||
if side == self.active_side:
|
if ready or self._active_knee_retracting():
|
||||||
self.phase = DeadBugPhase.NEED_RESET
|
self.phase = DeadBugPhase.NEED_RESET
|
||||||
|
self._reset_frames = 1 if ready else 0
|
||||||
elif self.phase == DeadBugPhase.NEED_RESET:
|
elif self.phase == DeadBugPhase.NEED_RESET:
|
||||||
if ready:
|
if ready:
|
||||||
self._reset_frames += 1
|
self._reset_frames += 1
|
||||||
@@ -51,21 +79,106 @@ class DeadBugStateMachine:
|
|||||||
self._reset_frames = 0
|
self._reset_frames = 0
|
||||||
|
|
||||||
feedback = list(metrics.feedback)
|
feedback = list(metrics.feedback)
|
||||||
if side is None and not ready:
|
display_side = detect_diagonal_extension(metrics)
|
||||||
|
if display_side is None and not ready:
|
||||||
feedback.append("Extend opposite arm and leg only")
|
feedback.append("Extend opposite arm and leg only")
|
||||||
if ready:
|
if ready:
|
||||||
feedback.append("Ready position")
|
feedback.append("Ready position")
|
||||||
elif side == "left_arm_right_leg":
|
elif display_side == "left_arm_right_leg":
|
||||||
feedback.append("Left arm + right leg")
|
feedback.append("Left arm + right leg")
|
||||||
elif side == "right_arm_left_leg":
|
elif display_side == "right_arm_left_leg":
|
||||||
feedback.append("Right arm + left leg")
|
feedback.append("Right arm + left leg")
|
||||||
|
|
||||||
is_standard = side is not None and not metrics.feedback
|
is_standard = display_side is not None and not metrics.feedback
|
||||||
return DeadBugResult(
|
return DeadBugResult(
|
||||||
rep_count=self.rep_count,
|
rep_count=self.rep_count,
|
||||||
phase=self.phase,
|
phase=self.phase,
|
||||||
side=side,
|
side=display_side,
|
||||||
is_standard=is_standard,
|
is_standard=is_standard,
|
||||||
feedback=feedback[:3],
|
feedback=feedback[:3],
|
||||||
metrics=metrics,
|
metrics=metrics,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _update_knee_trends(self, metrics: DeadBugMetrics) -> None:
|
||||||
|
"""更新平滑膝角和本帧变化量,用连续趋势抵消单帧抖动。"""
|
||||||
|
previous_left = self._smooth_left_knee_angle
|
||||||
|
previous_right = self._smooth_right_knee_angle
|
||||||
|
|
||||||
|
if previous_left is None:
|
||||||
|
self._smooth_left_knee_angle = metrics.left_knee_angle
|
||||||
|
self._left_knee_delta = 0.0
|
||||||
|
else:
|
||||||
|
self._smooth_left_knee_angle = (
|
||||||
|
_SMOOTHING_ALPHA * metrics.left_knee_angle
|
||||||
|
+ (1.0 - _SMOOTHING_ALPHA) * previous_left
|
||||||
|
)
|
||||||
|
self._left_knee_delta = self._smooth_left_knee_angle - previous_left
|
||||||
|
|
||||||
|
if previous_right is None:
|
||||||
|
self._smooth_right_knee_angle = metrics.right_knee_angle
|
||||||
|
self._right_knee_delta = 0.0
|
||||||
|
else:
|
||||||
|
self._smooth_right_knee_angle = (
|
||||||
|
_SMOOTHING_ALPHA * metrics.right_knee_angle
|
||||||
|
+ (1.0 - _SMOOTHING_ALPHA) * previous_right
|
||||||
|
)
|
||||||
|
self._right_knee_delta = self._smooth_right_knee_angle - previous_right
|
||||||
|
|
||||||
|
def _detect_motion_side(self, metrics: DeadBugMetrics) -> str | None:
|
||||||
|
"""基于膝角趋势推断正在伸展的对角侧,手臂只作为辅助校验。"""
|
||||||
|
raw_side = detect_diagonal_extension(metrics)
|
||||||
|
if raw_side is not None and self._side_has_extension_motion(raw_side):
|
||||||
|
return raw_side
|
||||||
|
|
||||||
|
left = self._smooth_left_knee_angle
|
||||||
|
right = self._smooth_right_knee_angle
|
||||||
|
if left is None or right is None:
|
||||||
|
return raw_side
|
||||||
|
|
||||||
|
both_legs_high = left >= _EXTENSION_PEAK_ANGLE and right >= _EXTENSION_PEAK_ANGLE
|
||||||
|
if both_legs_high and abs(left - right) < _SIDE_ANGLE_MARGIN:
|
||||||
|
return None
|
||||||
|
|
||||||
|
if (
|
||||||
|
right >= _EXTENSION_START_ANGLE
|
||||||
|
and right - left >= _SIDE_ANGLE_MARGIN
|
||||||
|
and (self._right_knee_delta >= _TREND_DELTA_DEGREES or right >= _EXTENSION_PEAK_ANGLE)
|
||||||
|
):
|
||||||
|
return "left_arm_right_leg"
|
||||||
|
if (
|
||||||
|
left >= _EXTENSION_START_ANGLE
|
||||||
|
and left - right >= _SIDE_ANGLE_MARGIN
|
||||||
|
and (self._left_knee_delta >= _TREND_DELTA_DEGREES or left >= _EXTENSION_PEAK_ANGLE)
|
||||||
|
):
|
||||||
|
return "right_arm_left_leg"
|
||||||
|
return raw_side
|
||||||
|
|
||||||
|
def _side_has_extension_motion(self, side: str) -> bool:
|
||||||
|
"""确认对应腿处于伸展区或仍在伸展趋势中。"""
|
||||||
|
if side == "left_arm_right_leg":
|
||||||
|
angle = self._smooth_right_knee_angle
|
||||||
|
delta = self._right_knee_delta
|
||||||
|
else:
|
||||||
|
angle = self._smooth_left_knee_angle
|
||||||
|
delta = self._left_knee_delta
|
||||||
|
if angle is None:
|
||||||
|
return True
|
||||||
|
return angle >= _EXTENSION_PEAK_ANGLE or (
|
||||||
|
angle >= _EXTENSION_START_ANGLE and delta >= _TREND_DELTA_DEGREES
|
||||||
|
)
|
||||||
|
|
||||||
|
def _is_stable_ready(self, metrics: DeadBugMetrics) -> bool:
|
||||||
|
"""准备位需要双腿回到屈膝区域;使用平滑膝角避免单帧阈值跳变。"""
|
||||||
|
left = self._smooth_left_knee_angle
|
||||||
|
right = self._smooth_right_knee_angle
|
||||||
|
if left is None or right is None:
|
||||||
|
return is_ready_position(metrics)
|
||||||
|
return left <= _READY_KNEE_ANGLE and right <= _READY_KNEE_ANGLE
|
||||||
|
|
||||||
|
def _active_knee_retracting(self) -> bool:
|
||||||
|
"""确认已伸展的那条腿开始回收。"""
|
||||||
|
if self.active_side == "left_arm_right_leg":
|
||||||
|
return self._right_knee_delta <= -_TREND_DELTA_DEGREES
|
||||||
|
if self.active_side == "right_arm_left_leg":
|
||||||
|
return self._left_knee_delta <= -_TREND_DELTA_DEGREES
|
||||||
|
return False
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ from enum import Enum
|
|||||||
|
|
||||||
|
|
||||||
class DeadBugPhase(str, Enum):
|
class DeadBugPhase(str, Enum):
|
||||||
|
"""死虫式动作阶段枚举"""
|
||||||
READY = "ready"
|
READY = "ready"
|
||||||
EXTENDING = "extending"
|
EXTENDING = "extending"
|
||||||
NEED_RESET = "need_reset"
|
NEED_RESET = "need_reset"
|
||||||
@@ -13,6 +14,7 @@ class DeadBugPhase(str, Enum):
|
|||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
class Point:
|
class Point:
|
||||||
|
"""三维关键点坐标及可见度"""
|
||||||
x: float
|
x: float
|
||||||
y: float
|
y: float
|
||||||
z: float
|
z: float
|
||||||
@@ -21,6 +23,7 @@ class Point:
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class DeadBugMetrics:
|
class DeadBugMetrics:
|
||||||
|
"""四肢关节度量数据"""
|
||||||
left_arm_extended: bool
|
left_arm_extended: bool
|
||||||
right_arm_extended: bool
|
right_arm_extended: bool
|
||||||
left_leg_extended: bool
|
left_leg_extended: bool
|
||||||
@@ -34,6 +37,7 @@ class DeadBugMetrics:
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class DeadBugResult:
|
class DeadBugResult:
|
||||||
|
"""单帧检测结果"""
|
||||||
rep_count: int
|
rep_count: int
|
||||||
phase: DeadBugPhase
|
phase: DeadBugPhase
|
||||||
side: str | None
|
side: str | None
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ from app.signaling.websocket_server import main as serve
|
|||||||
|
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
|
"""应用入口:启动服务并运行WebSocket信令服务器"""
|
||||||
startup()
|
startup()
|
||||||
logger.info("Starting server...")
|
logger.info("Starting server...")
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ from app.exercises.dead_bug.types import DeadBugResult
|
|||||||
|
|
||||||
|
|
||||||
def draw_status_overlay(image: np.ndarray, result: DeadBugResult) -> None:
|
def draw_status_overlay(image: np.ndarray, result: DeadBugResult) -> None:
|
||||||
|
"""在图像上叠加动作状态信息(次数、阶段、反馈)"""
|
||||||
color = (60, 220, 90) if result.is_standard else (50, 180, 255)
|
color = (60, 220, 90) if result.is_standard else (50, 180, 255)
|
||||||
cv2.rectangle(image, (12, 12), (520, 142), (20, 20, 20), -1)
|
cv2.rectangle(image, (12, 12), (520, 142), (20, 20, 20), -1)
|
||||||
cv2.putText(image, f"Dead bug reps: {result.rep_count}", (28, 48), cv2.FONT_HERSHEY_SIMPLEX, 0.9, color, 2)
|
cv2.putText(image, f"Dead bug reps: {result.rep_count}", (28, 48), cv2.FONT_HERSHEY_SIMPLEX, 0.9, color, 2)
|
||||||
|
|||||||
@@ -18,6 +18,7 @@ def draw_landmarks(
|
|||||||
line_thickness: int = 2,
|
line_thickness: int = 2,
|
||||||
point_radius: int = 4,
|
point_radius: int = 4,
|
||||||
) -> None:
|
) -> None:
|
||||||
|
"""绘制人体骨架关键点与连接线(仅绘制可见度达标的点)"""
|
||||||
if connections is None:
|
if connections is None:
|
||||||
connections = _POSE_CONNECTIONS
|
connections = _POSE_CONNECTIONS
|
||||||
|
|
||||||
|
|||||||
@@ -6,16 +6,20 @@ WINDOW_NAME = "Android Camera (WebRTC)"
|
|||||||
|
|
||||||
|
|
||||||
def show_frame(image, window_name: str = WINDOW_NAME) -> None:
|
def show_frame(image, window_name: str = WINDOW_NAME) -> None:
|
||||||
|
"""在OpenCV窗口中显示图像帧"""
|
||||||
cv2.imshow(window_name, image)
|
cv2.imshow(window_name, image)
|
||||||
|
|
||||||
|
|
||||||
def wait_key(delay_ms: int = 1) -> int:
|
def wait_key(delay_ms: int = 1) -> int:
|
||||||
|
"""等待按键并返回ASCII码"""
|
||||||
return cv2.waitKey(delay_ms) & 0xFF
|
return cv2.waitKey(delay_ms) & 0xFF
|
||||||
|
|
||||||
|
|
||||||
def is_esc_pressed() -> bool:
|
def is_esc_pressed() -> bool:
|
||||||
|
"""检测ESC键是否被按下"""
|
||||||
return wait_key(1) == 27
|
return wait_key(1) == 27
|
||||||
|
|
||||||
|
|
||||||
def close_window() -> None:
|
def close_window() -> None:
|
||||||
|
"""关闭所有OpenCV窗口"""
|
||||||
cv2.destroyAllWindows()
|
cv2.destroyAllWindows()
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ from aiortc import RTCIceCandidate
|
|||||||
|
|
||||||
|
|
||||||
def parse_ice(data: dict[str, Any]) -> RTCIceCandidate | None:
|
def parse_ice(data: dict[str, Any]) -> RTCIceCandidate | None:
|
||||||
|
"""解析ICE候选者字符串为RTCIceCandidate对象"""
|
||||||
match = re.match(
|
match = re.match(
|
||||||
r'candidate:(\S+) (\d) (\S+) (\d+) (\S+) (\d+) typ (\S+)(?: raddr (\S+) rport (\d+))?',
|
r'candidate:(\S+) (\d) (\S+) (\d+) (\S+) (\d+) typ (\S+)(?: raddr (\S+) rport (\d+))?',
|
||||||
data["candidate"],
|
data["candidate"],
|
||||||
|
|||||||
@@ -1,10 +1,11 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class SignalingMessage:
|
class SignalingMessage:
|
||||||
|
"""WebRTC信令消息数据模型"""
|
||||||
type: str
|
type: str
|
||||||
sdp: str = ""
|
sdp: str = ""
|
||||||
candidate: str = ""
|
candidate: str = ""
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ from configs.load import config
|
|||||||
|
|
||||||
|
|
||||||
async def handle_client(websocket):
|
async def handle_client(websocket):
|
||||||
|
"""处理单个WebSocket客户端连接"""
|
||||||
client = websocket.remote_address
|
client = websocket.remote_address
|
||||||
logger.info(f"Client connected: {client}")
|
logger.info(f"Client connected: {client}")
|
||||||
|
|
||||||
@@ -21,6 +22,7 @@ async def handle_client(websocket):
|
|||||||
|
|
||||||
|
|
||||||
async def main():
|
async def main():
|
||||||
|
"""启动WebSocket信令服务器"""
|
||||||
cfg = config.server
|
cfg = config.server
|
||||||
logger.info(f"WebRTC signaling server: ws://{cfg.host}:{cfg.port}")
|
logger.info(f"WebRTC signaling server: ws://{cfg.host}:{cfg.port}")
|
||||||
async with websockets.serve(handle_client, cfg.host, cfg.port, max_size=cfg.max_ws_size):
|
async with websockets.serve(handle_client, cfg.host, cfg.port, max_size=cfg.max_ws_size):
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ TARGET_HEIGHT = 720
|
|||||||
|
|
||||||
|
|
||||||
def resize_to_target(image: np.ndarray, width: int = TARGET_WIDTH, height: int = TARGET_HEIGHT) -> np.ndarray:
|
def resize_to_target(image: np.ndarray, width: int = TARGET_WIDTH, height: int = TARGET_HEIGHT) -> np.ndarray:
|
||||||
|
"""将图像缩放到目标尺寸(仅当尺寸不一致时)"""
|
||||||
h, w = image.shape[:2]
|
h, w = image.shape[:2]
|
||||||
if w == width and h == height:
|
if w == width and h == height:
|
||||||
return image
|
return image
|
||||||
@@ -16,8 +17,10 @@ def resize_to_target(image: np.ndarray, width: int = TARGET_WIDTH, height: int =
|
|||||||
|
|
||||||
|
|
||||||
def bgr_to_rgba(bgr: np.ndarray) -> np.ndarray:
|
def bgr_to_rgba(bgr: np.ndarray) -> np.ndarray:
|
||||||
|
"""将BGR格式图像转换为RGBA格式"""
|
||||||
return cv2.cvtColor(bgr, cv2.COLOR_BGR2RGBA)
|
return cv2.cvtColor(bgr, cv2.COLOR_BGR2RGBA)
|
||||||
|
|
||||||
|
|
||||||
def bgr_to_rgb(bgr: np.ndarray) -> np.ndarray:
|
def bgr_to_rgb(bgr: np.ndarray) -> np.ndarray:
|
||||||
|
"""将BGR格式图像转换为RGB格式"""
|
||||||
return cv2.cvtColor(bgr, cv2.COLOR_BGR2RGB)
|
return cv2.cvtColor(bgr, cv2.COLOR_BGR2RGB)
|
||||||
|
|||||||
@@ -1,7 +1,6 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import threading
|
import platform
|
||||||
import time
|
|
||||||
from typing import Callable
|
from typing import Callable
|
||||||
|
|
||||||
import mediapipe as mp
|
import mediapipe as mp
|
||||||
@@ -14,8 +13,9 @@ PoseLandmarkerOptions = mp.tasks.vision.PoseLandmarkerOptions
|
|||||||
VisionRunningMode = mp.tasks.vision.RunningMode
|
VisionRunningMode = mp.tasks.vision.RunningMode
|
||||||
BaseOptions = mp.tasks.BaseOptions
|
BaseOptions = mp.tasks.BaseOptions
|
||||||
|
|
||||||
|
|
||||||
class PoseLandmarkerWrapper:
|
class PoseLandmarkerWrapper:
|
||||||
|
"""MediaPipe姿态关键点检测器封装"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -23,22 +23,29 @@ class PoseLandmarkerWrapper:
|
|||||||
prefer_gpu: bool = True,
|
prefer_gpu: bool = True,
|
||||||
result_callback: Callable | None = None,
|
result_callback: Callable | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
|
"""初始化姿态检测器,优先尝试GPU委托,失败则回退到CPU"""
|
||||||
self.model_path = model_path or DEFAULT_MODEL_PATH
|
self.model_path = model_path or DEFAULT_MODEL_PATH
|
||||||
|
|
||||||
if prefer_gpu:
|
if prefer_gpu:
|
||||||
try:
|
try:
|
||||||
|
if platform.system() == "Windows":
|
||||||
|
logger.warning(
|
||||||
|
"MediaPipe GPU delegate requested, but MediaPipe Tasks Python does not support GPU delegate on Windows; "
|
||||||
|
"Intel Iris Xe cannot be used by this backend and CPU fallback is expected"
|
||||||
|
)
|
||||||
self.delegate = BaseOptions.Delegate.GPU
|
self.delegate = BaseOptions.Delegate.GPU
|
||||||
self._landmarker = self._create(PoseLandmarker.Delegate.GPU)
|
self._landmarker = self._create(self.delegate, result_callback)
|
||||||
logger.info("MediaPipe PoseLandmarker initialized with GPU delegate")
|
logger.info("MediaPipe PoseLandmarker initialized with GPU delegate")
|
||||||
return
|
return
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.warning("MediaPipe GPU delegate unavailable, falling back to CPU: {}", exc)
|
logger.warning("MediaPipe GPU delegate unavailable, falling back to CPU: {}", exc)
|
||||||
|
|
||||||
self.delegate = BaseOptions.Delegate.CPU
|
self.delegate = BaseOptions.Delegate.CPU
|
||||||
self._landmarker = self._create(PoseLandmarker.Delegate.CPU, result_callback)
|
self._landmarker = self._create(self.delegate, result_callback)
|
||||||
logger.info("MediaPipe PoseLandmarker initialized with CPU delegate")
|
logger.info("MediaPipe PoseLandmarker initialized with CPU delegate")
|
||||||
|
|
||||||
def _create(self, delegate, result_callback=None):
|
def _create(self, delegate, result_callback=None):
|
||||||
|
"""根据委托类型和回调创建PoseLandmarker实例"""
|
||||||
options = PoseLandmarkerOptions(
|
options = PoseLandmarkerOptions(
|
||||||
base_options=BaseOptions(model_asset_path=self.model_path, delegate=delegate),
|
base_options=BaseOptions(model_asset_path=self.model_path, delegate=delegate),
|
||||||
running_mode=VisionRunningMode.LIVE_STREAM,
|
running_mode=VisionRunningMode.LIVE_STREAM,
|
||||||
@@ -51,7 +58,9 @@ class PoseLandmarkerWrapper:
|
|||||||
return PoseLandmarker.create_from_options(options)
|
return PoseLandmarker.create_from_options(options)
|
||||||
|
|
||||||
def detect_async(self, mp_image, timestamp_ms: int) -> None:
|
def detect_async(self, mp_image, timestamp_ms: int) -> None:
|
||||||
|
"""异步执行姿态检测"""
|
||||||
return self._landmarker.detect_async(mp_image, timestamp_ms)
|
return self._landmarker.detect_async(mp_image, timestamp_ms)
|
||||||
|
|
||||||
def close(self) -> None:
|
def close(self) -> None:
|
||||||
|
"""释放MediaPipe资源"""
|
||||||
self._landmarker.close()
|
self._landmarker.close()
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ TARGET_HEIGHT = 720
|
|||||||
|
|
||||||
|
|
||||||
def validate_frame_size(image: np.ndarray, width: int = TARGET_WIDTH, height: int = TARGET_HEIGHT) -> None:
|
def validate_frame_size(image: np.ndarray, width: int = TARGET_WIDTH, height: int = TARGET_HEIGHT) -> None:
|
||||||
|
"""验证视频帧尺寸是否与目标尺寸一致,不一致时记录警告"""
|
||||||
h, w = image.shape[:2]
|
h, w = image.shape[:2]
|
||||||
if w != width or h != height:
|
if w != width or h != height:
|
||||||
logger.warning("Unexpected frame size: {}x{}", w, h)
|
logger.warning("Unexpected frame size: {}x{}", w, h)
|
||||||
|
|||||||
@@ -10,13 +10,15 @@ from loguru import logger
|
|||||||
from app.signaling.ice_parser import parse_ice
|
from app.signaling.ice_parser import parse_ice
|
||||||
from app.webrtc.video_receiver import VideoReceiver
|
from app.webrtc.video_receiver import VideoReceiver
|
||||||
|
|
||||||
|
|
||||||
class PeerSession:
|
class PeerSession:
|
||||||
|
"""WebRTC对等连接会话管理"""
|
||||||
|
|
||||||
def __init__(self) -> None:
|
def __init__(self) -> None:
|
||||||
self._pc = RTCPeerConnection()
|
self._pc = RTCPeerConnection()
|
||||||
self._video_task: asyncio.Task | None = None
|
self._video_task: asyncio.Task | None = None
|
||||||
|
|
||||||
async def handle(self, websocket) -> None:
|
async def handle(self, websocket) -> None:
|
||||||
|
"""处理WebSocket信令交互与WebRTC连接建立"""
|
||||||
self._setup_events()
|
self._setup_events()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -47,6 +49,7 @@ class PeerSession:
|
|||||||
await self._cleanup()
|
await self._cleanup()
|
||||||
|
|
||||||
def _setup_events(self) -> None:
|
def _setup_events(self) -> None:
|
||||||
|
"""注册ICE连接状态变化和视频轨道接收事件处理器"""
|
||||||
@self._pc.on("track")
|
@self._pc.on("track")
|
||||||
async def on_track(track):
|
async def on_track(track):
|
||||||
logger.info(f"Track received: kind={track.kind}")
|
logger.info(f"Track received: kind={track.kind}")
|
||||||
@@ -61,6 +64,7 @@ class PeerSession:
|
|||||||
await self._pc.close()
|
await self._pc.close()
|
||||||
|
|
||||||
async def _cleanup(self) -> None:
|
async def _cleanup(self) -> None:
|
||||||
|
"""清理视频任务并关闭对等连接"""
|
||||||
if self._video_task:
|
if self._video_task:
|
||||||
self._video_task.cancel()
|
self._video_task.cancel()
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import os
|
import time
|
||||||
|
|
||||||
import cv2
|
import cv2
|
||||||
from aiortc.mediastreams import MediaStreamError
|
from aiortc.mediastreams import MediaStreamError
|
||||||
@@ -14,6 +14,7 @@ from configs.load import config
|
|||||||
|
|
||||||
|
|
||||||
def _format_pose_debug(pose_result) -> str:
|
def _format_pose_debug(pose_result) -> str:
|
||||||
|
"""格式化姿态检测结果用于调试日志输出"""
|
||||||
metrics = pose_result.metrics
|
metrics = pose_result.metrics
|
||||||
if metrics is None:
|
if metrics is None:
|
||||||
return "metrics=None"
|
return "metrics=None"
|
||||||
@@ -26,12 +27,61 @@ def _format_pose_debug(pose_result) -> str:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _new_perf_window() -> dict:
|
||||||
|
return {
|
||||||
|
"frames": 0,
|
||||||
|
"processed": 0,
|
||||||
|
"loop_ms": 0.0,
|
||||||
|
"to_ndarray_ms": 0.0,
|
||||||
|
"detect_ms": 0.0,
|
||||||
|
"show_ms": 0.0,
|
||||||
|
"max_loop_ms": 0.0,
|
||||||
|
"max_detect_ms": 0.0,
|
||||||
|
"detector": {},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _add_detector_timing(perf: dict, timing: dict[str, float | bool]) -> None:
|
||||||
|
detector = perf["detector"]
|
||||||
|
for key, value in timing.items():
|
||||||
|
if key == "submitted":
|
||||||
|
detector[key] = detector.get(key, 0) + (1 if value else 0)
|
||||||
|
continue
|
||||||
|
value = float(value)
|
||||||
|
detector[key] = detector.get(key, 0.0) + value
|
||||||
|
max_key = f"max_{key}"
|
||||||
|
detector[max_key] = max(detector.get(max_key, 0.0), value)
|
||||||
|
|
||||||
|
|
||||||
|
def _avg(perf: dict, key: str, denominator: int) -> float:
|
||||||
|
if denominator <= 0:
|
||||||
|
return 0.0
|
||||||
|
return perf.get(key, 0.0) / denominator
|
||||||
|
|
||||||
|
|
||||||
class VideoReceiver:
|
class VideoReceiver:
|
||||||
|
"""视频轨道接收与运动检测流水线"""
|
||||||
|
|
||||||
def __init__(self, track) -> None:
|
def __init__(self, track) -> None:
|
||||||
self._track = track
|
self._track = track
|
||||||
|
|
||||||
async def run(self) -> None:
|
async def run(self) -> None:
|
||||||
logger.info("Start receiving video frames, process_every_n={}", config.video.process_every_n_frames)
|
"""持续接收视频帧并进行姿态检测、渲染和语音播报"""
|
||||||
|
log_every_n_frames = max(1, config.video.log_every_n_frames)
|
||||||
|
perf_log_every_n_frames = max(1, config.video.perf_log_every_n_frames)
|
||||||
|
slow_frame_ms = max(0.0, config.video.slow_frame_ms)
|
||||||
|
logger.info(
|
||||||
|
"Start receiving video frames, process_every_n={}, log_every_n={}, perf_log_every_n={}, slow_frame_ms={}",
|
||||||
|
config.video.process_every_n_frames,
|
||||||
|
log_every_n_frames,
|
||||||
|
perf_log_every_n_frames,
|
||||||
|
slow_frame_ms,
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
"OpenCV OpenCL status: have_opencl={}, use_opencl={}",
|
||||||
|
cv2.ocl.haveOpenCL(),
|
||||||
|
cv2.ocl.useOpenCL(),
|
||||||
|
)
|
||||||
|
|
||||||
frame_count = 0
|
frame_count = 0
|
||||||
processed_count = 0
|
processed_count = 0
|
||||||
@@ -44,31 +94,51 @@ class VideoReceiver:
|
|||||||
)
|
)
|
||||||
announcer = RepAnnouncer(
|
announcer = RepAnnouncer(
|
||||||
enabled=config.audio.rep_announcer_enabled,
|
enabled=config.audio.rep_announcer_enabled,
|
||||||
rate=config.audio.rep_announcer_rate,
|
max_count=config.audio.rep_max_count,
|
||||||
volume=config.audio.rep_announcer_volume,
|
audio_dir=config.audio.resolved_audio_dir,
|
||||||
)
|
)
|
||||||
last_announced_rep = 0
|
last_announced_rep = 0
|
||||||
last_pose_result = None
|
last_pose_result = None
|
||||||
last_annotated = None
|
last_annotated = None
|
||||||
|
perf = _new_perf_window()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
while True:
|
while True:
|
||||||
|
loop_started = time.perf_counter()
|
||||||
frame = await self._track.recv()
|
frame = await self._track.recv()
|
||||||
frame_count += 1
|
frame_count += 1
|
||||||
|
recv_done = time.perf_counter()
|
||||||
raw_img = frame.to_ndarray(format="bgr24")
|
raw_img = frame.to_ndarray(format="bgr24")
|
||||||
|
ndarray_done = time.perf_counter()
|
||||||
timestamp_ms = int(frame.time * 1000) if frame.time is not None else frame_count * 33
|
timestamp_ms = int(frame.time * 1000) if frame.time is not None else frame_count * 33
|
||||||
|
|
||||||
|
detect_ms = 0.0
|
||||||
if frame_count % config.video.process_every_n_frames == 0 or last_pose_result is None:
|
if frame_count % config.video.process_every_n_frames == 0 or last_pose_result is None:
|
||||||
|
detect_started = time.perf_counter()
|
||||||
processed_count += 1
|
processed_count += 1
|
||||||
last_annotated, last_pose_result = detector.process_frame(raw_img, timestamp_ms)
|
last_annotated, last_pose_result = detector.process_frame(raw_img, timestamp_ms)
|
||||||
|
detect_ms = (time.perf_counter() - detect_started) * 1000
|
||||||
|
perf["processed"] += 1
|
||||||
|
perf["detect_ms"] += detect_ms
|
||||||
|
perf["max_detect_ms"] = max(perf["max_detect_ms"], detect_ms)
|
||||||
|
_add_detector_timing(perf, detector.last_timing)
|
||||||
if last_pose_result.rep_count > last_announced_rep:
|
if last_pose_result.rep_count > last_announced_rep:
|
||||||
last_announced_rep = last_pose_result.rep_count
|
last_announced_rep = last_pose_result.rep_count
|
||||||
|
announce_started = time.perf_counter()
|
||||||
announcer.announce_count(last_announced_rep)
|
announcer.announce_count(last_announced_rep)
|
||||||
|
logger.info(
|
||||||
|
"Rep completed and audio requested: count={}, frame={}, announce_call_ms={:.1f}",
|
||||||
|
last_announced_rep,
|
||||||
|
frame_count,
|
||||||
|
(time.perf_counter() - announce_started) * 1000,
|
||||||
|
)
|
||||||
|
|
||||||
display_img = last_annotated if last_annotated is not None else raw_img
|
display_img = last_annotated if last_annotated is not None else raw_img
|
||||||
|
show_started = time.perf_counter()
|
||||||
show_frame(display_img)
|
show_frame(display_img)
|
||||||
|
show_done = time.perf_counter()
|
||||||
|
|
||||||
if frame_count % 100 == 0:
|
if frame_count % log_every_n_frames == 0:
|
||||||
logger.info(
|
logger.info(
|
||||||
"Received {} frames, processed={}, raw_shape={}, reps={}, phase={}, feedback={}, {}",
|
"Received {} frames, processed={}, raw_shape={}, reps={}, phase={}, feedback={}, {}",
|
||||||
frame_count,
|
frame_count,
|
||||||
@@ -80,6 +150,53 @@ class VideoReceiver:
|
|||||||
_format_pose_debug(last_pose_result) if last_pose_result is not None else "metrics=None",
|
_format_pose_debug(last_pose_result) if last_pose_result is not None else "metrics=None",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
loop_ms = (show_done - loop_started) * 1000
|
||||||
|
to_ndarray_ms = (ndarray_done - recv_done) * 1000
|
||||||
|
show_ms = (show_done - show_started) * 1000
|
||||||
|
perf["frames"] += 1
|
||||||
|
perf["loop_ms"] += loop_ms
|
||||||
|
perf["to_ndarray_ms"] += to_ndarray_ms
|
||||||
|
perf["show_ms"] += show_ms
|
||||||
|
perf["max_loop_ms"] = max(perf["max_loop_ms"], loop_ms)
|
||||||
|
|
||||||
|
if slow_frame_ms and loop_ms >= slow_frame_ms:
|
||||||
|
logger.warning(
|
||||||
|
"Slow video frame: frame={}, loop_ms={:.1f}, detect_ms={:.1f}, to_ndarray_ms={:.1f}, show_ms={:.1f}, shape={}",
|
||||||
|
frame_count,
|
||||||
|
loop_ms,
|
||||||
|
detect_ms,
|
||||||
|
to_ndarray_ms,
|
||||||
|
show_ms,
|
||||||
|
raw_img.shape,
|
||||||
|
)
|
||||||
|
|
||||||
|
if frame_count % perf_log_every_n_frames == 0:
|
||||||
|
frames = perf["frames"]
|
||||||
|
processed = perf["processed"]
|
||||||
|
detector_perf = perf["detector"]
|
||||||
|
logger.info(
|
||||||
|
"Perf window: frames={}, processed={}, avg_loop_ms={:.1f}, max_loop_ms={:.1f}, avg_to_ndarray_ms={:.1f}, "
|
||||||
|
"avg_detect_ms={:.1f}, max_detect_ms={:.1f}, avg_show_ms={:.1f}, detector_avg_total_ms={:.1f}, "
|
||||||
|
"detector_max_total_ms={:.1f}, detector_avg_wait_ms={:.1f}, detector_max_wait_ms={:.1f}, "
|
||||||
|
"detector_avg_convert_ms={:.1f}, detector_avg_postprocess_draw_ms={:.1f}, detector_submitted={}",
|
||||||
|
frames,
|
||||||
|
processed,
|
||||||
|
_avg(perf, "loop_ms", frames),
|
||||||
|
perf["max_loop_ms"],
|
||||||
|
_avg(perf, "to_ndarray_ms", frames),
|
||||||
|
_avg(perf, "detect_ms", processed),
|
||||||
|
perf["max_detect_ms"],
|
||||||
|
_avg(perf, "show_ms", frames),
|
||||||
|
_avg(detector_perf, "total_ms", processed),
|
||||||
|
detector_perf.get("max_total_ms", 0.0),
|
||||||
|
_avg(detector_perf, "wait_ms", processed),
|
||||||
|
detector_perf.get("max_wait_ms", 0.0),
|
||||||
|
_avg(detector_perf, "convert_ms", processed),
|
||||||
|
_avg(detector_perf, "postprocess_draw_ms", processed),
|
||||||
|
detector_perf.get("submitted", 0),
|
||||||
|
)
|
||||||
|
perf = _new_perf_window()
|
||||||
|
|
||||||
if is_esc_pressed():
|
if is_esc_pressed():
|
||||||
logger.info("ESC pressed, closing display")
|
logger.info("ESC pressed, closing display")
|
||||||
break
|
break
|
||||||
|
|||||||
+10
-2
@@ -7,10 +7,13 @@ server:
|
|||||||
max_ws_size: 10485760 # 10 MB
|
max_ws_size: 10485760 # 10 MB
|
||||||
|
|
||||||
video:
|
video:
|
||||||
process_every_n_frames: 1
|
process_every_n_frames: 2
|
||||||
|
log_every_n_frames: 30
|
||||||
|
perf_log_every_n_frames: 30
|
||||||
|
slow_frame_ms: 100
|
||||||
|
|
||||||
model:
|
model:
|
||||||
path: "./pose_models/pose_landmarker_full.task" # empty = auto-detect pose_models/pose_landmarker_full.task
|
path: "./pose_models/pose_landmarker_full.task"
|
||||||
prefer_gpu: true
|
prefer_gpu: true
|
||||||
|
|
||||||
dead_bug:
|
dead_bug:
|
||||||
@@ -22,6 +25,11 @@ audio:
|
|||||||
rep_announcer_enabled: true
|
rep_announcer_enabled: true
|
||||||
rep_announcer_rate: 185
|
rep_announcer_rate: 185
|
||||||
rep_announcer_volume: 1.0
|
rep_announcer_volume: 1.0
|
||||||
|
rep_max_count: 200 # 预生成语音文件的最大次数
|
||||||
|
rep_audio_dir: "" # 空 = 自动使用 app/audio/reps
|
||||||
|
trim_leading_silence: true
|
||||||
|
trim_silence_threshold: 500
|
||||||
|
trim_silence_padding_ms: 20
|
||||||
|
|
||||||
logging:
|
logging:
|
||||||
dir: logs
|
dir: logs
|
||||||
|
|||||||
+3
-30
@@ -1,7 +1,6 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import dataclasses
|
import dataclasses
|
||||||
import os
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
@@ -19,24 +18,6 @@ from configs.models import (
|
|||||||
|
|
||||||
_PROJECT_ROOT = Path(__file__).resolve().parent.parent
|
_PROJECT_ROOT = Path(__file__).resolve().parent.parent
|
||||||
|
|
||||||
_ENV_MAP = {
|
|
||||||
"POSEFIT_WS_HOST": ("server", "host"),
|
|
||||||
"POSEFIT_WS_PORT": ("server", "port", int),
|
|
||||||
"POSEFIT_WS_MAX_SIZE": ("server", "max_ws_size", int),
|
|
||||||
"POSEFIT_PROCESS_EVERY_N_FRAMES": ("video", "process_every_n_frames", int),
|
|
||||||
"POSEFIT_MODEL_PATH": ("model", "path"),
|
|
||||||
"POSEFIT_PREFER_GPU": ("model", "prefer_gpu", lambda v: v not in ("0", "false", "False")),
|
|
||||||
"POSEFIT_VISIBILITY_THRESHOLD": ("dead_bug", "visibility_threshold", float),
|
|
||||||
"POSEFIT_EXTENSION_CONFIRM_FRAMES": ("dead_bug", "extension_confirm_frames", int),
|
|
||||||
"POSEFIT_RESET_CONFIRM_FRAMES": ("dead_bug", "reset_confirm_frames", int),
|
|
||||||
"POSEFIT_REP_ANNOUNCER_ENABLED": ("audio", "rep_announcer_enabled", lambda v: v not in ("0", "false", "False")),
|
|
||||||
"POSEFIT_REP_ANNOUNCER_RATE": ("audio", "rep_announcer_rate", int),
|
|
||||||
"POSEFIT_REP_ANNOUNCER_VOLUME": ("audio", "rep_announcer_volume", float),
|
|
||||||
"POSEFIT_LOG_ROTATION": ("logging", "rotation"),
|
|
||||||
"POSEFIT_LOG_RETENTION": ("logging", "retention"),
|
|
||||||
"POSEFIT_LOG_DIR": ("logging", "dir"),
|
|
||||||
}
|
|
||||||
|
|
||||||
_SECTION_CLASS = {
|
_SECTION_CLASS = {
|
||||||
"server": ServerConfig,
|
"server": ServerConfig,
|
||||||
"video": VideoConfig,
|
"video": VideoConfig,
|
||||||
@@ -48,6 +29,7 @@ _SECTION_CLASS = {
|
|||||||
|
|
||||||
|
|
||||||
def _dict_to_dataclass(cls: type, data: dict[str, Any] | None) -> dict[str, Any]:
|
def _dict_to_dataclass(cls: type, data: dict[str, Any] | None) -> dict[str, Any]:
|
||||||
|
"""将字典过滤为仅包含指定dataclass字段的键值对"""
|
||||||
if data is None:
|
if data is None:
|
||||||
return {}
|
return {}
|
||||||
field_names = {f.name for f in dataclasses.fields(cls)}
|
field_names = {f.name for f in dataclasses.fields(cls)}
|
||||||
@@ -55,28 +37,19 @@ def _dict_to_dataclass(cls: type, data: dict[str, Any] | None) -> dict[str, Any]
|
|||||||
|
|
||||||
|
|
||||||
def _read_yaml(path: Path) -> dict[str, Any]:
|
def _read_yaml(path: Path) -> dict[str, Any]:
|
||||||
|
"""读取YAML配置文件并返回字典"""
|
||||||
if not path.exists():
|
if not path.exists():
|
||||||
return {}
|
return {}
|
||||||
with open(path, encoding="utf-8") as f:
|
with open(path, encoding="utf-8") as f:
|
||||||
return yaml.safe_load(f) or {}
|
return yaml.safe_load(f) or {}
|
||||||
|
|
||||||
|
|
||||||
def _apply_env_overrides(raw: dict[str, Any]) -> None:
|
|
||||||
for env_var, (section, key, *rest) in _ENV_MAP.items():
|
|
||||||
value = os.getenv(env_var)
|
|
||||||
if value is None:
|
|
||||||
continue
|
|
||||||
if rest:
|
|
||||||
value = rest[0](value)
|
|
||||||
raw.setdefault(section, {})[key] = value
|
|
||||||
|
|
||||||
|
|
||||||
def load_config(config_path: str | Path | None = None) -> AppConfig:
|
def load_config(config_path: str | Path | None = None) -> AppConfig:
|
||||||
|
"""加载并解析应用配置,返回AppConfig实例"""
|
||||||
if config_path is None:
|
if config_path is None:
|
||||||
config_path = _PROJECT_ROOT / "config.yaml"
|
config_path = _PROJECT_ROOT / "config.yaml"
|
||||||
|
|
||||||
raw = _read_yaml(Path(config_path))
|
raw = _read_yaml(Path(config_path))
|
||||||
_apply_env_overrides(raw)
|
|
||||||
|
|
||||||
return AppConfig(**{
|
return AppConfig(**{
|
||||||
section: cls(**_dict_to_dataclass(cls, raw.get(section)))
|
section: cls(**_dict_to_dataclass(cls, raw.get(section)))
|
||||||
|
|||||||
+25
-1
@@ -6,6 +6,7 @@ from pathlib import Path
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ServerConfig:
|
class ServerConfig:
|
||||||
|
"""WebSocket服务器配置"""
|
||||||
host: str = "0.0.0.0"
|
host: str = "0.0.0.0"
|
||||||
port: int = 8765
|
port: int = 8765
|
||||||
max_ws_size: int = 10_485_760
|
max_ws_size: int = 10_485_760
|
||||||
@@ -13,16 +14,22 @@ class ServerConfig:
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class VideoConfig:
|
class VideoConfig:
|
||||||
process_every_n_frames: int = 1
|
"""视频帧处理配置"""
|
||||||
|
process_every_n_frames: int = 2
|
||||||
|
log_every_n_frames: int = 30
|
||||||
|
perf_log_every_n_frames: int = 30
|
||||||
|
slow_frame_ms: float = 100.0
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ModelConfig:
|
class ModelConfig:
|
||||||
|
"""姿态检测模型配置"""
|
||||||
path: str = ""
|
path: str = ""
|
||||||
prefer_gpu: bool = True
|
prefer_gpu: bool = True
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def resolved_path(self) -> str:
|
def resolved_path(self) -> str:
|
||||||
|
"""返回模型文件的绝对路径"""
|
||||||
if self.path:
|
if self.path:
|
||||||
return self.path
|
return self.path
|
||||||
return str(Path(__file__).resolve().parent.parent / "pose_models" / "pose_landmarker_full.task")
|
return str(Path(__file__).resolve().parent.parent / "pose_models" / "pose_landmarker_full.task")
|
||||||
@@ -30,6 +37,7 @@ class ModelConfig:
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class DeadBugConfig:
|
class DeadBugConfig:
|
||||||
|
"""死虫式(Dead Bug)运动检测配置"""
|
||||||
visibility_threshold: float = 0.45
|
visibility_threshold: float = 0.45
|
||||||
extension_confirm_frames: int = 4
|
extension_confirm_frames: int = 4
|
||||||
reset_confirm_frames: int = 3
|
reset_confirm_frames: int = 3
|
||||||
@@ -37,24 +45,40 @@ class DeadBugConfig:
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class AudioConfig:
|
class AudioConfig:
|
||||||
|
"""语音播报配置"""
|
||||||
rep_announcer_enabled: bool = True
|
rep_announcer_enabled: bool = True
|
||||||
rep_announcer_rate: int = 185
|
rep_announcer_rate: int = 185
|
||||||
rep_announcer_volume: float = 1.0
|
rep_announcer_volume: float = 1.0
|
||||||
|
rep_max_count: int = 200
|
||||||
|
rep_audio_dir: str = ""
|
||||||
|
trim_leading_silence: bool = True
|
||||||
|
trim_silence_threshold: int = 500
|
||||||
|
trim_silence_padding_ms: int = 20
|
||||||
|
|
||||||
|
@property
|
||||||
|
def resolved_audio_dir(self) -> Path:
|
||||||
|
"""返回语音文件目录的绝对路径"""
|
||||||
|
if self.rep_audio_dir:
|
||||||
|
return Path(self.rep_audio_dir)
|
||||||
|
return Path(__file__).resolve().parent.parent / "resources" / "audio" / "reps"
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class LoggingConfig:
|
class LoggingConfig:
|
||||||
|
"""日志配置"""
|
||||||
dir: str = "logs"
|
dir: str = "logs"
|
||||||
rotation: str = "20 MB"
|
rotation: str = "20 MB"
|
||||||
retention: str = "14 days"
|
retention: str = "14 days"
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def dir_path(self) -> Path:
|
def dir_path(self) -> Path:
|
||||||
|
"""返回日志目录的绝对路径"""
|
||||||
return Path(__file__).resolve().parent.parent / self.dir
|
return Path(__file__).resolve().parent.parent / self.dir
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class AppConfig:
|
class AppConfig:
|
||||||
|
"""应用总配置,聚合所有子配置"""
|
||||||
server: ServerConfig = field(default_factory=ServerConfig)
|
server: ServerConfig = field(default_factory=ServerConfig)
|
||||||
video: VideoConfig = field(default_factory=VideoConfig)
|
video: VideoConfig = field(default_factory=VideoConfig)
|
||||||
model: ModelConfig = field(default_factory=ModelConfig)
|
model: ModelConfig = field(default_factory=ModelConfig)
|
||||||
|
|||||||
@@ -3,26 +3,32 @@ from __future__ import annotations
|
|||||||
from app.exercises.dead_bug.rules import detect_diagonal_extension, has_required_visibility, is_ready_position
|
from app.exercises.dead_bug.rules import detect_diagonal_extension, has_required_visibility, is_ready_position
|
||||||
from app.exercises.dead_bug.types import DeadBugMetrics, Point
|
from app.exercises.dead_bug.types import DeadBugMetrics, Point
|
||||||
|
|
||||||
|
|
||||||
class TestDeadBugRules:
|
class TestDeadBugRules:
|
||||||
|
"""死虫式规则函数单元测试"""
|
||||||
|
|
||||||
def _make_landmark(self, x=0.5, y=0.5, z=0.0, visibility=1.0):
|
def _make_landmark(self, x=0.5, y=0.5, z=0.0, visibility=1.0):
|
||||||
|
"""创建测试用Point对象"""
|
||||||
return Point(x, y, z, visibility)
|
return Point(x, y, z, visibility)
|
||||||
|
|
||||||
def _make_visible_landmarks(self):
|
def _make_visible_landmarks(self):
|
||||||
|
"""创建33个全可见的测试用关键点"""
|
||||||
return [self._make_landmark() for _ in range(33)]
|
return [self._make_landmark() for _ in range(33)]
|
||||||
|
|
||||||
def test_has_required_visibility_all_visible(self):
|
def test_has_required_visibility_all_visible(self):
|
||||||
|
"""测试:所有关键点可见时应返回True"""
|
||||||
lm = self._make_visible_landmarks()
|
lm = self._make_visible_landmarks()
|
||||||
indices = (11, 12, 13, 14, 15, 16, 23, 24, 25, 26, 27, 28)
|
indices = (11, 12, 13, 14, 15, 16, 23, 24, 25, 26, 27, 28)
|
||||||
assert has_required_visibility(lm, indices, 0.45)
|
assert has_required_visibility(lm, indices, 0.45)
|
||||||
|
|
||||||
def test_has_required_visibility_low(self):
|
def test_has_required_visibility_low(self):
|
||||||
|
"""测试:关键点可见度过低时应返回False"""
|
||||||
lm = self._make_visible_landmarks()
|
lm = self._make_visible_landmarks()
|
||||||
lm[11] = self._make_landmark(visibility=0.1)
|
lm[11] = self._make_landmark(visibility=0.1)
|
||||||
indices = (11, 12, 13, 14, 15, 16, 23, 24, 25, 26, 27, 28)
|
indices = (11, 12, 13, 14, 15, 16, 23, 24, 25, 26, 27, 28)
|
||||||
assert not has_required_visibility(lm, indices, 0.45)
|
assert not has_required_visibility(lm, indices, 0.45)
|
||||||
|
|
||||||
def test_detect_diagonal_extension_none(self):
|
def test_detect_diagonal_extension_none(self):
|
||||||
|
"""测试:四肢均未伸展时应返回None"""
|
||||||
metrics = DeadBugMetrics(
|
metrics = DeadBugMetrics(
|
||||||
left_arm_extended=False, right_arm_extended=False,
|
left_arm_extended=False, right_arm_extended=False,
|
||||||
left_leg_extended=False, right_leg_extended=False,
|
left_leg_extended=False, right_leg_extended=False,
|
||||||
@@ -33,6 +39,7 @@ class TestDeadBugRules:
|
|||||||
assert detect_diagonal_extension(metrics) is None
|
assert detect_diagonal_extension(metrics) is None
|
||||||
|
|
||||||
def test_detect_diagonal_extension_left_arm_right_leg(self):
|
def test_detect_diagonal_extension_left_arm_right_leg(self):
|
||||||
|
"""测试:左臂+右腿对角伸展检测"""
|
||||||
metrics = DeadBugMetrics(
|
metrics = DeadBugMetrics(
|
||||||
left_arm_extended=True, right_arm_extended=False,
|
left_arm_extended=True, right_arm_extended=False,
|
||||||
left_leg_extended=False, right_leg_extended=True,
|
left_leg_extended=False, right_leg_extended=True,
|
||||||
@@ -42,7 +49,30 @@ class TestDeadBugRules:
|
|||||||
)
|
)
|
||||||
assert detect_diagonal_extension(metrics) == "left_arm_right_leg"
|
assert detect_diagonal_extension(metrics) == "left_arm_right_leg"
|
||||||
|
|
||||||
|
def test_detect_diagonal_extension_allows_ready_arm_overlap(self):
|
||||||
|
"""测试:准备位双臂上举时,单腿对侧伸展仍应识别"""
|
||||||
|
metrics = DeadBugMetrics(
|
||||||
|
left_arm_extended=True, right_arm_extended=True,
|
||||||
|
left_leg_extended=False, right_leg_extended=True,
|
||||||
|
left_elbow_angle=160, right_elbow_angle=160,
|
||||||
|
left_knee_angle=90, right_knee_angle=160,
|
||||||
|
feedback=[],
|
||||||
|
)
|
||||||
|
assert detect_diagonal_extension(metrics) == "left_arm_right_leg"
|
||||||
|
|
||||||
|
def test_detect_diagonal_extension_rejects_both_legs(self):
|
||||||
|
"""测试:双腿同时伸展不应识别为可计数对角伸展"""
|
||||||
|
metrics = DeadBugMetrics(
|
||||||
|
left_arm_extended=True, right_arm_extended=True,
|
||||||
|
left_leg_extended=True, right_leg_extended=True,
|
||||||
|
left_elbow_angle=160, right_elbow_angle=160,
|
||||||
|
left_knee_angle=160, right_knee_angle=160,
|
||||||
|
feedback=[],
|
||||||
|
)
|
||||||
|
assert detect_diagonal_extension(metrics) is None
|
||||||
|
|
||||||
def test_is_ready_position(self):
|
def test_is_ready_position(self):
|
||||||
|
"""测试:膝盖弯曲且四肢收缩应识别为准备姿态"""
|
||||||
metrics = DeadBugMetrics(
|
metrics = DeadBugMetrics(
|
||||||
left_arm_extended=False, right_arm_extended=False,
|
left_arm_extended=False, right_arm_extended=False,
|
||||||
left_leg_extended=False, right_leg_extended=False,
|
left_leg_extended=False, right_leg_extended=False,
|
||||||
@@ -52,7 +82,19 @@ class TestDeadBugRules:
|
|||||||
)
|
)
|
||||||
assert is_ready_position(metrics)
|
assert is_ready_position(metrics)
|
||||||
|
|
||||||
|
def test_is_ready_allows_arms_extended(self):
|
||||||
|
"""测试:dead bug 准备位允许双臂上举"""
|
||||||
|
metrics = DeadBugMetrics(
|
||||||
|
left_arm_extended=True, right_arm_extended=True,
|
||||||
|
left_leg_extended=False, right_leg_extended=False,
|
||||||
|
left_elbow_angle=160, right_elbow_angle=160,
|
||||||
|
left_knee_angle=100, right_knee_angle=100,
|
||||||
|
feedback=[],
|
||||||
|
)
|
||||||
|
assert is_ready_position(metrics)
|
||||||
|
|
||||||
def test_is_not_ready_legs_extended(self):
|
def test_is_not_ready_legs_extended(self):
|
||||||
|
"""测试:腿部伸展时不识别为准备姿态"""
|
||||||
metrics = DeadBugMetrics(
|
metrics = DeadBugMetrics(
|
||||||
left_arm_extended=False, right_arm_extended=False,
|
left_arm_extended=False, right_arm_extended=False,
|
||||||
left_leg_extended=True, right_leg_extended=False,
|
left_leg_extended=True, right_leg_extended=False,
|
||||||
|
|||||||
@@ -3,9 +3,11 @@ from __future__ import annotations
|
|||||||
from app.exercises.dead_bug.state_machine import DeadBugStateMachine
|
from app.exercises.dead_bug.state_machine import DeadBugStateMachine
|
||||||
from app.exercises.dead_bug.types import DeadBugMetrics, DeadBugPhase
|
from app.exercises.dead_bug.types import DeadBugMetrics, DeadBugPhase
|
||||||
|
|
||||||
|
|
||||||
class TestDeadBugStateMachine:
|
class TestDeadBugStateMachine:
|
||||||
|
"""死虫式状态机单元测试"""
|
||||||
|
|
||||||
def _ready_metrics(self) -> DeadBugMetrics:
|
def _ready_metrics(self) -> DeadBugMetrics:
|
||||||
|
"""构建准备姿态的度量数据"""
|
||||||
return DeadBugMetrics(
|
return DeadBugMetrics(
|
||||||
left_arm_extended=False, right_arm_extended=False,
|
left_arm_extended=False, right_arm_extended=False,
|
||||||
left_leg_extended=False, right_leg_extended=False,
|
left_leg_extended=False, right_leg_extended=False,
|
||||||
@@ -15,6 +17,7 @@ class TestDeadBugStateMachine:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def _extended_left(self) -> DeadBugMetrics:
|
def _extended_left(self) -> DeadBugMetrics:
|
||||||
|
"""构建左臂+右腿对角伸展的度量数据"""
|
||||||
return DeadBugMetrics(
|
return DeadBugMetrics(
|
||||||
left_arm_extended=True, right_arm_extended=False,
|
left_arm_extended=True, right_arm_extended=False,
|
||||||
left_leg_extended=False, right_leg_extended=True,
|
left_leg_extended=False, right_leg_extended=True,
|
||||||
@@ -23,20 +26,108 @@ class TestDeadBugStateMachine:
|
|||||||
feedback=[],
|
feedback=[],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _both_legs_extended(self) -> DeadBugMetrics:
|
||||||
|
"""构建双腿同时伸展的非标准姿态"""
|
||||||
|
return DeadBugMetrics(
|
||||||
|
left_arm_extended=True, right_arm_extended=True,
|
||||||
|
left_leg_extended=True, right_leg_extended=True,
|
||||||
|
left_elbow_angle=160, right_elbow_angle=160,
|
||||||
|
left_knee_angle=160, right_knee_angle=160,
|
||||||
|
feedback=[],
|
||||||
|
)
|
||||||
|
|
||||||
|
def _arms_extended_ready_legs(self) -> DeadBugMetrics:
|
||||||
|
"""构建腿已收回但手臂未收回的姿态"""
|
||||||
|
return DeadBugMetrics(
|
||||||
|
left_arm_extended=True, right_arm_extended=True,
|
||||||
|
left_leg_extended=False, right_leg_extended=False,
|
||||||
|
left_elbow_angle=160, right_elbow_angle=160,
|
||||||
|
left_knee_angle=100, right_knee_angle=100,
|
||||||
|
feedback=[],
|
||||||
|
)
|
||||||
|
|
||||||
|
def _right_knee_angle(self, angle: float) -> DeadBugMetrics:
|
||||||
|
"""构建右膝角连续变化但伸展布尔值尚未稳定的姿态"""
|
||||||
|
return DeadBugMetrics(
|
||||||
|
left_arm_extended=True, right_arm_extended=True,
|
||||||
|
left_leg_extended=False, right_leg_extended=False,
|
||||||
|
left_elbow_angle=160, right_elbow_angle=160,
|
||||||
|
left_knee_angle=90, right_knee_angle=angle,
|
||||||
|
feedback=[],
|
||||||
|
)
|
||||||
|
|
||||||
def test_initial_state(self):
|
def test_initial_state(self):
|
||||||
|
"""测试:状态机初始化后应为READY且计数为0"""
|
||||||
sm = DeadBugStateMachine()
|
sm = DeadBugStateMachine()
|
||||||
assert sm.phase == DeadBugPhase.READY
|
assert sm.phase == DeadBugPhase.READY
|
||||||
assert sm.rep_count == 0
|
assert sm.rep_count == 0
|
||||||
|
|
||||||
def test_no_transition_in_ready(self):
|
def test_no_transition_in_ready(self):
|
||||||
|
"""测试:准备姿态下不触发状态转换"""
|
||||||
sm = DeadBugStateMachine()
|
sm = DeadBugStateMachine()
|
||||||
result = sm.update(self._ready_metrics())
|
result = sm.update(self._ready_metrics())
|
||||||
assert sm.phase == DeadBugPhase.READY
|
assert sm.phase == DeadBugPhase.READY
|
||||||
assert result.rep_count == 0
|
assert result.rep_count == 0
|
||||||
|
|
||||||
def test_confirm_extension(self):
|
def test_confirm_extension(self):
|
||||||
|
"""测试:连续确认帧数后从READY转换到EXTENDING"""
|
||||||
sm = DeadBugStateMachine(extension_confirm_frames=2, reset_confirm_frames=2)
|
sm = DeadBugStateMachine(extension_confirm_frames=2, reset_confirm_frames=2)
|
||||||
sm.update(self._extended_left())
|
sm.update(self._extended_left())
|
||||||
assert sm.phase == DeadBugPhase.READY
|
assert sm.phase == DeadBugPhase.READY
|
||||||
sm.update(self._extended_left())
|
sm.update(self._extended_left())
|
||||||
assert sm.phase == DeadBugPhase.EXTENDING
|
assert sm.phase == DeadBugPhase.EXTENDING
|
||||||
|
|
||||||
|
def test_confirm_extension_from_knee_angle_trend(self):
|
||||||
|
"""测试:膝角连续上升时,不依赖单帧伸展布尔值也能确认伸展"""
|
||||||
|
sm = DeadBugStateMachine(extension_confirm_frames=2, reset_confirm_frames=2)
|
||||||
|
|
||||||
|
sm.update(self._right_knee_angle(100))
|
||||||
|
sm.update(self._right_knee_angle(130))
|
||||||
|
assert sm.phase == DeadBugPhase.READY
|
||||||
|
sm.update(self._right_knee_angle(145))
|
||||||
|
sm.update(self._right_knee_angle(150))
|
||||||
|
|
||||||
|
assert sm.phase == DeadBugPhase.EXTENDING
|
||||||
|
|
||||||
|
def test_full_rep_counts_once_after_strict_reset(self):
|
||||||
|
"""测试:确认伸展后,只有严格回到准备姿态才计一次"""
|
||||||
|
sm = DeadBugStateMachine(extension_confirm_frames=2, reset_confirm_frames=2)
|
||||||
|
|
||||||
|
sm.update(self._extended_left())
|
||||||
|
sm.update(self._extended_left())
|
||||||
|
assert sm.phase == DeadBugPhase.EXTENDING
|
||||||
|
|
||||||
|
sm.update(self._arms_extended_ready_legs())
|
||||||
|
assert sm.rep_count == 0
|
||||||
|
assert sm.phase == DeadBugPhase.NEED_RESET
|
||||||
|
|
||||||
|
sm.update(self._ready_metrics())
|
||||||
|
result = sm.update(self._ready_metrics())
|
||||||
|
assert result.rep_count == 1
|
||||||
|
assert sm.phase == DeadBugPhase.READY
|
||||||
|
|
||||||
|
result = sm.update(self._ready_metrics())
|
||||||
|
assert result.rep_count == 1
|
||||||
|
|
||||||
|
def test_both_legs_do_not_start_rep(self):
|
||||||
|
"""测试:双腿同时伸展不进入计数流程"""
|
||||||
|
sm = DeadBugStateMachine(extension_confirm_frames=2, reset_confirm_frames=2)
|
||||||
|
sm.update(self._both_legs_extended())
|
||||||
|
result = sm.update(self._both_legs_extended())
|
||||||
|
|
||||||
|
assert result.rep_count == 0
|
||||||
|
assert sm.phase == DeadBugPhase.READY
|
||||||
|
|
||||||
|
def test_no_pose_preserves_confirmed_rep_until_reset(self):
|
||||||
|
"""测试:已确认伸展后短暂丢姿态,回到准备位仍能完成计数"""
|
||||||
|
sm = DeadBugStateMachine(extension_confirm_frames=2, reset_confirm_frames=2)
|
||||||
|
sm.update(self._extended_left())
|
||||||
|
sm.update(self._extended_left())
|
||||||
|
assert sm.phase == DeadBugPhase.EXTENDING
|
||||||
|
|
||||||
|
sm.mark_no_pose()
|
||||||
|
sm.update(self._ready_metrics())
|
||||||
|
result = sm.update(self._ready_metrics())
|
||||||
|
|
||||||
|
assert result.rep_count == 1
|
||||||
|
assert sm.phase == DeadBugPhase.READY
|
||||||
|
|||||||
@@ -2,9 +2,11 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from app.signaling.ice_parser import parse_ice
|
from app.signaling.ice_parser import parse_ice
|
||||||
|
|
||||||
|
|
||||||
class TestIceParser:
|
class TestIceParser:
|
||||||
|
"""ICE候选者解析单元测试"""
|
||||||
|
|
||||||
def test_parse_valid_ice(self):
|
def test_parse_valid_ice(self):
|
||||||
|
"""测试:解析有效的ICE host候选者"""
|
||||||
data = {
|
data = {
|
||||||
"candidate": "1234567890 1 UDP 2130706431 192.168.1.1 12345 typ host",
|
"candidate": "1234567890 1 UDP 2130706431 192.168.1.1 12345 typ host",
|
||||||
"sdpMid": "0",
|
"sdpMid": "0",
|
||||||
@@ -20,9 +22,11 @@ class TestIceParser:
|
|||||||
assert cand.type == "host"
|
assert cand.type == "host"
|
||||||
|
|
||||||
def test_parse_invalid_ice(self):
|
def test_parse_invalid_ice(self):
|
||||||
|
"""测试:解析无效ICE字符串应返回None"""
|
||||||
assert parse_ice({"candidate": "invalid"}) is None
|
assert parse_ice({"candidate": "invalid"}) is None
|
||||||
|
|
||||||
def test_parse_srflx(self):
|
def test_parse_srflx(self):
|
||||||
|
"""测试:解析含有raddr/rport的srflx候选者"""
|
||||||
data = {
|
data = {
|
||||||
"candidate": "abcdef 1 UDP 1686052607 203.0.113.1 50000 typ srflx raddr 192.168.1.1 rport 12345",
|
"candidate": "abcdef 1 UDP 1686052607 203.0.113.1 50000 typ srflx raddr 192.168.1.1 rport 12345",
|
||||||
"sdpMid": "0",
|
"sdpMid": "0",
|
||||||
|
|||||||
Reference in New Issue
Block a user