Files
product/scripts/smoke_test.py
T
2026-03-27 10:49:34 +08:00

104 lines
3.4 KiB
Python

from __future__ import annotations
import argparse
import asyncio
import math
import os
import sys
import wave
import cv2
import numpy as np
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
if PROJECT_ROOT not in sys.path:
sys.path.insert(0, PROJECT_ROOT)
from core.pipeline import ChatPipeline
from core.state_machine import Arbitrator
from services.asr import ASRService
from services.avatar import AvatarService
from services.llm import LLMService
from services.tts import TTSService
def _write_wav(path: str, audio: np.ndarray, sr: int) -> None:
pcm = np.clip(audio * 32767.0, -32768, 32767).astype(np.int16)
with wave.open(path, "wb") as wf:
wf.setnchannels(1)
wf.setsampwidth(2)
wf.setframerate(sr)
wf.writeframes(pcm.tobytes())
async def _write_avatar_video(
avatar: AvatarService, audio: np.ndarray, sr: int, out_path: str, fps: int
) -> tuple[int, float]:
fourcc = cv2.VideoWriter_fourcc(*"mp4v")
writer = cv2.VideoWriter(out_path, fourcc, float(fps), (avatar.width, avatar.height))
if not writer.isOpened():
raise RuntimeError(f"failed to open video writer: {out_path}")
duration_s = float(audio.shape[0]) / float(sr) if sr > 0 else 0.0
total_frames = max(1, int(math.ceil(duration_s * fps)))
samples_per_frame = max(1, int(sr / fps))
for i in range(total_frames):
s = i * samples_per_frame
e = min(audio.shape[0], (i + 1) * samples_per_frame)
seg = audio[s:e]
if seg.size == 0:
level = 0.0
else:
rms = float(np.sqrt(np.mean(seg.astype(np.float32) ** 2)))
level = float(np.clip(rms * 10.0, 0.0, 1.0))
frame = await avatar.render_frame(mouth_open=level, speaking=True)
writer.write(frame)
# add 0.4s idle tail for easy visual check
tail_frames = max(1, int(0.4 * fps))
for _ in range(tail_frames):
frame = await avatar.render_frame(mouth_open=0.0, speaking=False)
writer.write(frame)
writer.release()
return total_frames + tail_frames, duration_s + 0.4
async def main() -> None:
parser = argparse.ArgumentParser(description="Offline full-pipeline smoke test")
parser.add_argument("--text", default="你好,请做一个十秒内的简短自我介绍。")
parser.add_argument("--out-dir", default="outputs")
parser.add_argument("--fps", type=int, default=25)
args = parser.parse_args()
os.makedirs(args.out_dir, exist_ok=True)
wav_path = os.path.join(args.out_dir, "smoke_reply.wav")
mp4_path = os.path.join(args.out_dir, "smoke_avatar.mp4")
arbitrator = Arbitrator()
pipeline = ChatPipeline(
arbitrator=arbitrator,
asr=ASRService(),
llm=LLMService(),
tts=TTSService(),
)
avatar = AvatarService()
result, audio, sr = await pipeline.process_text_turn(args.text)
_write_wav(wav_path, audio, sr)
frame_count, video_s = await _write_avatar_video(avatar, audio, sr, mp4_path, args.fps)
print("=== smoke test done ===")
print(f"user_text: {result.get('user_text', '')}")
print(f"reply_text: {result.get('reply_text', '')}")
print(f"llm_source: {result.get('llm_source', 'unknown')}")
print(f"audio_samples: {result.get('audio_samples', 0)}, sample_rate: {sr}")
print(f"video_frames: {frame_count}, video_seconds: {video_s:.2f}")
print(f"wav: {wav_path}")
print(f"mp4: {mp4_path}")
if __name__ == "__main__":
asyncio.run(main())