Compare commits

..

12 Commits

Author SHA1 Message Date
wsy182 08b6543b79 perf(video): 优化视频处理性能监控和音频播放
- 添加视频处理性能计时和统计功能
- 实现帧处理时间监控和慢帧警告
- 添加音频文件静音修剪功能
- 优化Windows平台音频播放实现
- 调整默认日志输出频率减少冗余信息
- 修复MediaPipe GPU委托在Windows上的兼容性问题
2026-06-15 23:13:36 +08:00
wsy182 6dee2a2ff3 feat(exercise): 优化死虫式训练姿态检测算法
- 调整视频处理频率从每帧处理改为每2帧处理
- 添加膝角趋势平滑算法减少单帧抖动误判
- 改进对角伸展检测逻辑支持准备位手臂上举
- 优化状态机确保严格回到准备姿态才计数
- 添加姿态丢失时的候选帧清理机制
- 更新音频文件生成路径至resources目录
- 改进macOS音频生成使用AIFF格式提高质量
- 添加详细的帧处理日志输出间隔配置
2026-06-10 22:57:35 +08:00
wsy182 ea0c007441 Clean up RepAnnouncer: remove TTS code, play pre-generated audio only
- RepAnnouncer now only plays audio files, no TTS generation
- Removed pyttsx3 dependency, rate/volume params from constructor
- Audio generation delegated to app/audio/generate.py (called at startup)
- Default audio dir changed to resources/audio/reps
- Added resources/ to .gitignore
2026-06-10 11:51:05 +08:00
wsy182 b45a8e2e85 Add audio generation config, refactor rep_announcer
- AudioConfig now includes rep_max_count and rep_audio_dir
- app/audio/generate.py uses config instead of hardcoded constants
- RepAnnouncer rewrote with pre-generated audio cache
- Supports Windows winsound, macOS afplay, Linux paplay/aplay
- Pin requirements back to mediapipe==0.10.21 with numpy<2
2026-06-10 11:42:40 +08:00
wsy182 1f6c3f3de8 refactor(vision): 优化姿态关键点检测器的初始化逻辑
- 移除未使用的 threading 和 time 模块导入
- 统一委托类型的使用,避免硬编码委托类型
- 简化 GPU 和 CPU 委托的创建流程
- 修复委托类型传递的一致性问题
2026-06-10 11:26:39 +08:00
wsy182 37b85cd683 chore(deps): 更新依赖包版本
- 将 numpy 版本从 >=1.26,<2 更新为 >=2.4.6
- 将 mediapipe 版本从 ==0.10.21 更新为 ==0.10.35
- 保持其他依赖包版本不变
- 确保依赖版本兼容性
2026-06-10 10:35:45 +08:00
wsy182 c3f93e4441 为所有函数和类添加中文注释文档字符串 2026-06-10 10:34:11 +08:00
wsy182 c612a7ad71 style(config): 移除配置文件中的多余注释和环境变量设置
- 移除了 config.yaml 中的注释行
- 移除了 app/core/logging.py 中多余的空格
- 移除了 app/main.py 中的环境变量设置代码
- 移除了 app/main.py 中未使用的导入语句
2026-06-10 10:29:36 +08:00
wsy182 e86c2301ec Remove env var overrides, config.yaml is single source of truth
- Stripped _ENV_MAP and _apply_env_overrides from configs/load.py
- Cleaned up unused imports in video_receiver.py
- Restored MediaPipe env suppressors in app/main.py
- Removed .env.example (replaced by config.yaml)
2026-06-10 10:26:48 +08:00
wsy182 ae52578ed7 Separate config models from loader
- configs/models.py: AppConfig and all section dataclasses
- configs/load.py: pure loading logic (yaml, env overrides)
- config = load_config() singleton for consumers
2026-06-10 10:24:45 +08:00
wsy182 f9384f7bc1 Parse config.yaml into typed AppConfig dataclass
- ServerConfig, VideoConfig, ModelConfig, DeadBugConfig,
  AudioConfig, LoggingConfig as nested dataclasses
- Consumers use config.server.host, config.model.resolved_path etc.
- env var overrides preserved via _apply_env_overrides()
2026-06-10 10:23:51 +08:00
wsy182 c8fd057129 Centralize configuration into config.yaml
- All settings moved to config.yaml
- configs/load.py reads from config.yaml with env var overrides
- Environment variables still work for backward compatibility
- Added pyyaml to requirements
2026-06-10 10:19:41 +08:00
34 changed files with 1119 additions and 170 deletions
-24
View File
@@ -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
View File
@@ -2,3 +2,5 @@
.idea/ .idea/
__pycache__/ __pycache__/
logs/ logs/
resources/
+1 -5
View File
@@ -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.
+227
View File
@@ -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
View File
@@ -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
+15 -3
View File
@@ -1,10 +1,22 @@
from __future__ import annotations from __future__ import annotations
from app.diagnostics.crash_handler import enable_crash_handler from app.diagnostics.crash_handler import enable_crash_handler
from configs.default import LOG_DIR from configs.load import config
from app.audio.generate import generate_rep_audio_files
def startup() -> None: def startup() -> None:
enable_crash_handler(LOG_DIR) """应用启动初始化:开启崩溃日志和日志系统"""
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,
)
+7 -5
View File
@@ -4,16 +4,18 @@ from pathlib import Path
from loguru import logger from loguru import logger
from configs.default import LOG_DIR, LOG_RETENTION, LOG_ROTATION from configs.load import config
def setup_logging() -> None: def setup_logging() -> None:
LOG_DIR.mkdir(parents=True, exist_ok=True) """配置loguru日志输出到按日期轮转的日志文件"""
log_dir = config.logging.dir_path
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=LOG_ROTATION, rotation=config.logging.rotation,
retention=LOG_RETENTION, retention=config.logging.retention,
enqueue=True, enqueue=True,
backtrace=True, backtrace=True,
diagnose=True, diagnose=True,
+1
View File
@@ -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)
+6 -1
View File
@@ -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()
+83 -1
View File
@@ -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
+3
View File
@@ -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 -1
View File
@@ -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
+121 -8
View File
@@ -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
+4
View File
@@ -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
+13 -7
View File
@@ -1,10 +1,5 @@
from __future__ import annotations from __future__ import annotations
import os
os.environ["MEDIAPIPE_DISABLE_LOGGING"] = "1"
os.environ["GLOG_minloglevel"] = "3"
import asyncio import asyncio
from loguru import logger from loguru import logger
@@ -13,7 +8,18 @@ from app.core.lifecycle import startup
from app.signaling.websocket_server import main as serve from app.signaling.websocket_server import main as serve
if __name__ == "__main__": def main():
"""应用入口:启动服务并运行WebSocket信令服务器"""
startup() startup()
logger.info("Starting server...") logger.info("Starting server...")
asyncio.run(serve()) try:
asyncio.run(serve())
except (KeyboardInterrupt, SystemExit):
logger.info("Server stopped by user")
except Exception as e:
logger.error(f"Server error: {e}")
raise
if __name__ == "__main__":
main()
+1
View File
@@ -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)
+1
View File
@@ -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
+4
View File
@@ -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()
+1
View File
@@ -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"],
+2 -1
View File
@@ -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 = ""
+6 -3
View File
@@ -7,10 +7,11 @@ import websockets
from loguru import logger from loguru import logger
from app.webrtc.peer_session import PeerSession from app.webrtc.peer_session import PeerSession
from configs.default import WS_HOST, WS_MAX_SIZE, WS_PORT 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,8 @@ async def handle_client(websocket):
async def main(): async def main():
logger.info(f"WebRTC signaling server: ws://{WS_HOST}:{WS_PORT}") """启动WebSocket信令服务器"""
async with websockets.serve(handle_client, WS_HOST, WS_PORT, max_size=WS_MAX_SIZE): cfg = config.server
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):
await asyncio.Future() await asyncio.Future()
+3
View File
@@ -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)
+14 -5
View File
@@ -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()
+1
View File
@@ -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)
+5 -1
View File
@@ -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:
+130 -23
View File
@@ -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
@@ -10,20 +10,11 @@ from loguru import logger
from app.audio.rep_announcer import RepAnnouncer from app.audio.rep_announcer import RepAnnouncer
from app.exercises.dead_bug.detector import DeadBugDetector from app.exercises.dead_bug.detector import DeadBugDetector
from app.rendering.window_display import close_window, is_esc_pressed, show_frame from app.rendering.window_display import close_window, is_esc_pressed, show_frame
from configs.default import ( from configs.load import config
EXTENSION_CONFIRM_FRAMES,
MODEL_PATH,
PREFER_GPU,
PROCESS_EVERY_N_FRAMES,
REP_ANNOUNCER_ENABLED,
REP_ANNOUNCER_RATE,
REP_ANNOUNCER_VOLUME,
RESET_CONFIRM_FRAMES,
VISIBILITY_THRESHOLD,
)
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"
@@ -36,49 +27,118 @@ 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={}", 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
detector = DeadBugDetector( detector = DeadBugDetector(
model_path=MODEL_PATH, model_path=config.model.resolved_path,
visibility_threshold=VISIBILITY_THRESHOLD, visibility_threshold=config.dead_bug.visibility_threshold,
extension_confirm_frames=EXTENSION_CONFIRM_FRAMES, extension_confirm_frames=config.dead_bug.extension_confirm_frames,
reset_confirm_frames=RESET_CONFIRM_FRAMES, reset_confirm_frames=config.dead_bug.reset_confirm_frames,
prefer_gpu=PREFER_GPU, prefer_gpu=config.model.prefer_gpu,
) )
announcer = RepAnnouncer( announcer = RepAnnouncer(
enabled=REP_ANNOUNCER_ENABLED, enabled=config.audio.rep_announcer_enabled,
rate=REP_ANNOUNCER_RATE, max_count=config.audio.rep_max_count,
volume=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
if frame_count % PROCESS_EVERY_N_FRAMES == 0 or last_pose_result is None: detect_ms = 0.0
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,
@@ -90,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
+37
View File
@@ -0,0 +1,37 @@
# PoseFit Server Configuration
# ============================
server:
host: "0.0.0.0"
port: 8765
max_ws_size: 10485760 # 10 MB
video:
process_every_n_frames: 2
log_every_n_frames: 30
perf_log_every_n_frames: 30
slow_frame_ms: 100
model:
path: "./pose_models/pose_landmarker_full.task"
prefer_gpu: true
dead_bug:
visibility_threshold: 0.45
extension_confirm_frames: 4
reset_confirm_frames: 3
audio:
rep_announcer_enabled: true
rep_announcer_rate: 185
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:
dir: logs
rotation: "20 MB"
retention: "14 days"
-32
View File
@@ -1,32 +0,0 @@
from __future__ import annotations
import os
from pathlib import Path
# ── Server ──────────────────────────────────────────────────────────────────
WS_HOST = os.getenv("POSEFIT_WS_HOST", "0.0.0.0")
WS_PORT = int(os.getenv("POSEFIT_WS_PORT", "8765"))
WS_MAX_SIZE = int(os.getenv("POSEFIT_WS_MAX_SIZE", str(10 * 1024 * 1024)))
# ── Video processing ────────────────────────────────────────────────────────
PROCESS_EVERY_N_FRAMES = max(1, int(os.getenv("POSEFIT_PROCESS_EVERY_N_FRAMES", "1")))
# ── Model ───────────────────────────────────────────────────────────────────
MODEL_DIR: Path = Path(__file__).resolve().parent.parent / "pose_models"
MODEL_PATH = os.getenv("POSEFIT_MODEL_PATH", str(MODEL_DIR / "pose_landmarker_full.task"))
PREFER_GPU = os.getenv("POSEFIT_PREFER_GPU", "1") not in ("0", "false", "False")
# ── Dead bug exercise ───────────────────────────────────────────────────────
VISIBILITY_THRESHOLD = float(os.getenv("POSEFIT_VISIBILITY_THRESHOLD", "0.45"))
EXTENSION_CONFIRM_FRAMES = int(os.getenv("POSEFIT_EXTENSION_CONFIRM_FRAMES", "4"))
RESET_CONFIRM_FRAMES = int(os.getenv("POSEFIT_RESET_CONFIRM_FRAMES", "3"))
# ── Audio ───────────────────────────────────────────────────────────────────
REP_ANNOUNCER_ENABLED = os.getenv("POSEFIT_REP_ANNOUNCER_ENABLED", "1") not in ("0", "false", "False")
REP_ANNOUNCER_RATE = int(os.getenv("POSEFIT_REP_ANNOUNCER_RATE", "185"))
REP_ANNOUNCER_VOLUME = float(os.getenv("POSEFIT_REP_ANNOUNCER_VOLUME", "1.0"))
# ── Logging ─────────────────────────────────────────────────────────────────
LOG_DIR: Path = Path(__file__).resolve().parent.parent / "logs"
LOG_ROTATION = os.getenv("POSEFIT_LOG_ROTATION", "20 MB")
LOG_RETENTION = os.getenv("POSEFIT_LOG_RETENTION", "14 days")
+60
View File
@@ -0,0 +1,60 @@
from __future__ import annotations
import dataclasses
from pathlib import Path
from typing import Any
import yaml
from configs.models import (
AppConfig,
AudioConfig,
DeadBugConfig,
LoggingConfig,
ModelConfig,
ServerConfig,
VideoConfig,
)
_PROJECT_ROOT = Path(__file__).resolve().parent.parent
_SECTION_CLASS = {
"server": ServerConfig,
"video": VideoConfig,
"model": ModelConfig,
"dead_bug": DeadBugConfig,
"audio": AudioConfig,
"logging": LoggingConfig,
}
def _dict_to_dataclass(cls: type, data: dict[str, Any] | None) -> dict[str, Any]:
"""将字典过滤为仅包含指定dataclass字段的键值对"""
if data is None:
return {}
field_names = {f.name for f in dataclasses.fields(cls)}
return {k: v for k, v in data.items() if k in field_names}
def _read_yaml(path: Path) -> dict[str, Any]:
"""读取YAML配置文件并返回字典"""
if not path.exists():
return {}
with open(path, encoding="utf-8") as f:
return yaml.safe_load(f) or {}
def load_config(config_path: str | Path | None = None) -> AppConfig:
"""加载并解析应用配置,返回AppConfig实例"""
if config_path is None:
config_path = _PROJECT_ROOT / "config.yaml"
raw = _read_yaml(Path(config_path))
return AppConfig(**{
section: cls(**_dict_to_dataclass(cls, raw.get(section)))
for section, cls in _SECTION_CLASS.items()
})
config = load_config()
+87
View File
@@ -0,0 +1,87 @@
from __future__ import annotations
from dataclasses import dataclass, field
from pathlib import Path
@dataclass
class ServerConfig:
"""WebSocket服务器配置"""
host: str = "0.0.0.0"
port: int = 8765
max_ws_size: int = 10_485_760
@dataclass
class VideoConfig:
"""视频帧处理配置"""
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
class ModelConfig:
"""姿态检测模型配置"""
path: str = ""
prefer_gpu: bool = True
@property
def resolved_path(self) -> str:
"""返回模型文件的绝对路径"""
if self.path:
return self.path
return str(Path(__file__).resolve().parent.parent / "pose_models" / "pose_landmarker_full.task")
@dataclass
class DeadBugConfig:
"""死虫式(Dead Bug)运动检测配置"""
visibility_threshold: float = 0.45
extension_confirm_frames: int = 4
reset_confirm_frames: int = 3
@dataclass
class AudioConfig:
"""语音播报配置"""
rep_announcer_enabled: bool = True
rep_announcer_rate: int = 185
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
class LoggingConfig:
"""日志配置"""
dir: str = "logs"
rotation: str = "20 MB"
retention: str = "14 days"
@property
def dir_path(self) -> Path:
"""返回日志目录的绝对路径"""
return Path(__file__).resolve().parent.parent / self.dir
@dataclass
class AppConfig:
"""应用总配置,聚合所有子配置"""
server: ServerConfig = field(default_factory=ServerConfig)
video: VideoConfig = field(default_factory=VideoConfig)
model: ModelConfig = field(default_factory=ModelConfig)
dead_bug: DeadBugConfig = field(default_factory=DeadBugConfig)
audio: AudioConfig = field(default_factory=AudioConfig)
logging: LoggingConfig = field(default_factory=LoggingConfig)
+1
View File
@@ -5,3 +5,4 @@ numpy>=1.26,<2
loguru>=0.7.0 loguru>=0.7.0
mediapipe==0.10.21 mediapipe==0.10.21
pyttsx3>=2.99 pyttsx3>=2.99
pyyaml>=6.0
+44 -4
View File
@@ -1,30 +1,34 @@
from __future__ import annotations from __future__ import annotations
from app.exercises.dead_bug.metrics import calculate_metrics
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.state_machine import DeadBugStateMachine from app.exercises.dead_bug.types import DeadBugMetrics, Point
from app.exercises.dead_bug.types import DeadBugMetrics, DeadBugPhase, 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,
@@ -35,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,
@@ -44,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,
@@ -54,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,
+92 -1
View File
@@ -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
+5 -1
View File
@@ -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",