104 lines
3.4 KiB
Python
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())
|