#!/usr/bin/env python3 """ Multi-user concurrent streaming test for MuseTalk gRPC Avatar Service Uses DIRECT gRPC connection (no SSH tunnel required) Usage: python benchmarks/avatar_multi_user_test.py --server 81.166.173.12:10597 --audio data/audio/test.wav python benchmarks/avatar_multi_user_test.py --server localhost:50052 --users 1,2,3,5 """ import sys import os import time import wave import asyncio import argparse import statistics # Add grpc module path sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..', 'server', 'grpc')) import grpc from grpc import aio import avatar_pb2 import avatar_pb2_grpc async def single_user_stream(server: str, user_id: int, audio_data: bytes, sample_rate: int, chunk_ms: int = 100): """Run a streaming session for a single user with its own connection""" chunk_size = int(sample_rate * chunk_ms / 1000) * 2 session_id = f"user_{user_id}_{int(time.time()*1000)}" # Each user gets their own channel channel = aio.insecure_channel(server) stub = avatar_pb2_grpc.AvatarServiceStub(channel) async def audio_generator(): for i in range(0, len(audio_data), chunk_size): chunk = audio_data[i:i + chunk_size] is_final = (i + chunk_size >= len(audio_data)) yield avatar_pb2.AudioChunk( session_id=session_id, audio_data=chunk, sample_rate=sample_rate, is_final=is_final, timestamp=i / 2 / sample_rate, chunk_index=i // chunk_size ) # Simulate real-time streaming (0.3x to send faster) await asyncio.sleep(chunk_ms / 1000 * 0.3) start_time = time.time() first_frame_time = None frame_count = 0 total_bytes = 0 errors = [] try: async for response in stub.StreamingGenerate(audio_generator()): if first_frame_time is None: first_frame_time = time.time() frame_count += 1 total_bytes += len(response.frame_data) if response.is_final: break except Exception as e: errors.append(str(e)) finally: await channel.close() end_time = time.time() ttff = (first_frame_time - start_time) * 1000 if first_frame_time else None total_time = end_time - start_time stream_time = end_time - first_frame_time if first_frame_time else 0 fps = frame_count / stream_time if stream_time > 0 else 0 return { 'user_id': user_id, 'session_id': session_id, 'ttff_ms': ttff, 'total_time': total_time, 'frame_count': frame_count, 'fps': fps, 'total_bytes': total_bytes, 'errors': errors } async def run_concurrent_test(num_users: int, server: str, audio_path: str): """Run concurrent streaming test with multiple users""" print(f"\n{'='*60}") print(f"Concurrent Streaming Test: {num_users} users") print(f"{'='*60}") # Load audio with wave.open(audio_path, 'rb') as wf: sample_rate = wf.getframerate() n_frames = wf.getnframes() duration = n_frames / sample_rate audio_data = wf.readframes(n_frames) print(f"[Audio] Duration: {duration:.2f}s, Sample rate: {sample_rate}Hz") # Quick health check channel = aio.insecure_channel(server) stub = avatar_pb2_grpc.AvatarServiceStub(channel) try: health = await stub.HealthCheck(avatar_pb2.HealthRequest()) print(f"[Server] Status: {health.status}, Avatar: {health.avatar_id}") except Exception as e: print(f"[Server] Health check failed: {e}") await channel.close() return None info = await stub.GetAvatarInfo(avatar_pb2.AvatarInfoRequest()) print(f"[Avatar] {info.avatar_id}, {info.width}x{info.height}, {info.fps}fps") await channel.close() # Launch concurrent users (each with own connection) print(f"\n[Test] Starting {num_users} concurrent users...") start_all = time.time() tasks = [ single_user_stream(server, i, audio_data, sample_rate) for i in range(num_users) ] results = await asyncio.gather(*tasks, return_exceptions=True) end_all = time.time() # Analyze results successful = [r for r in results if isinstance(r, dict) and r['ttff_ms'] is not None] failed = [r for r in results if not isinstance(r, dict) or r['ttff_ms'] is None] print(f"\n[Results] {num_users} users") print(f" Successful: {len(successful)}/{num_users}") print(f" Failed: {len(failed)}/{num_users}") total_frames = 0 if successful: ttffs = [r['ttff_ms'] for r in successful] fps_vals = [r['fps'] for r in successful] frame_counts = [r['frame_count'] for r in successful] print(f"\n Time to First Frame (TTFF):") print(f" Min: {min(ttffs):.0f}ms") print(f" Max: {max(ttffs):.0f}ms") print(f" Avg: {statistics.mean(ttffs):.0f}ms") print(f" Median: {statistics.median(ttffs):.0f}ms") print(f"\n FPS per user:") print(f" Min: {min(fps_vals):.1f}") print(f" Max: {max(fps_vals):.1f}") print(f" Avg: {statistics.mean(fps_vals):.1f}") print(f"\n Frame counts:") print(f" Min: {min(frame_counts)}") print(f" Max: {max(frame_counts)}") total_frames = sum(frame_counts) print(f"\n Total throughput: {total_frames} frames in {end_all - start_all:.1f}s") print(f" Aggregate FPS: {total_frames / (end_all - start_all):.1f}") if failed: print(f"\n Failures:") for r in failed: if isinstance(r, Exception): print(f" Exception: {r}") else: print(f" User {r.get('user_id', '?')}: {r.get('errors', ['Unknown error'])}") return { 'num_users': num_users, 'successful': len(successful), 'failed': len(failed), 'avg_ttff_ms': statistics.mean(ttffs) if successful else None, 'max_ttff_ms': max(ttffs) if successful else None, 'avg_fps': statistics.mean(fps_vals) if successful else None, 'aggregate_fps': total_frames / (end_all - start_all) if successful else None } async def main(): parser = argparse.ArgumentParser(description='MuseTalk Avatar Multi-User Benchmark') parser.add_argument('--server', type=str, default='81.166.173.12:10597', help='gRPC server address (default: 81.166.173.12:10597)') parser.add_argument('--audio', type=str, default=os.path.join(os.path.dirname(__file__), '..', 'orpheus_demo_3_phrases.wav'), help='Path to audio file for testing') parser.add_argument('--users', type=str, default='1,2,3,5', help='Comma-separated list of user counts to test (default: 1,2,3,5)') parser.add_argument('--wait', type=int, default=5, help='Seconds to wait between tests (default: 5)') args = parser.parse_args() user_counts = [int(x.strip()) for x in args.users.split(',')] print("="*60) print("MuseTalk gRPC Multi-User Benchmark") print(f"Server: {args.server}") print(f"Audio: {args.audio}") print(f"User counts: {user_counts}") print("="*60) all_results = [] for num_users in user_counts: result = await run_concurrent_test(num_users, args.server, args.audio) if result: all_results.append(result) # Check if server is struggling if result['failed'] > 0 or (result['avg_ttff_ms'] and result['avg_ttff_ms'] > 10000): print(f"\n[WARNING] Server struggling at {num_users} users") if result['failed'] == num_users: print("[STOPPING] All users failed, stopping test") break # Give server time to recover if num_users < max(user_counts): print(f"\n[Wait] Waiting {args.wait}s before next test...") await asyncio.sleep(args.wait) # Summary print("\n" + "="*60) print("SUMMARY") print("="*60) print(f"{'Users':>6} | {'Success':>7} | {'TTFF (ms)':>10} | {'Avg FPS':>8} | {'Total FPS':>10}") print("-"*60) for r in all_results: ttff_str = f"{r['avg_ttff_ms']:.0f}" if r['avg_ttff_ms'] else "N/A" fps_str = f"{r['avg_fps']:.1f}" if r['avg_fps'] else "N/A" agg_str = f"{r['aggregate_fps']:.1f}" if r['aggregate_fps'] else "N/A" print(f"{r['num_users']:>6} | {r['successful']:>7} | {ttff_str:>10} | {fps_str:>8} | {agg_str:>10}") print("="*60) # Recommendation if all_results: best = max([r for r in all_results if r['failed'] == 0], key=lambda x: x['num_users'], default=None) if best: print(f"\nRecommended max users: {best['num_users']}") print(f" - All users successful") print(f" - Average TTFF: {best['avg_ttff_ms']:.0f}ms") print(f" - Average FPS per user: {best['avg_fps']:.1f}") if __name__ == '__main__': asyncio.run(main())