#!/usr/bin/env python3 """ Test script for BCT engine integration Based on rpc_engine_test.ipynb structure """ import os import time import argparse import multiprocessing as mp from tqdm import tqdm # Set CUDA device os.environ["CUDA_VISIBLE_DEVICES"] = "4" print("Setting up BCT engine test...") # Import required modules from suno_utils.gpt.generation import GenerationConfig from suno_utils.gpt.generation_engine import make_request from suno_utils.gpt.rpc_zmq_bct import start_bct_service_processes # Import tracing modules from suno_utils.tracing import start_daemon def test_bct_engine(num_requests=4, max_gen_duration_s=30): # Start trace daemon first print("Starting trace daemon...") daemon_proc = mp.Process( target=start_daemon, args=(18862, 20.0, 20.0, "s3", "bct_engine_test"), # port delay duration backend project ) daemon_proc.start() # Wait a moment for daemon to start time.sleep(2) print("Starting BCT service processes...") # Start BCT engine service client = start_bct_service_processes( "/app2/suno/checkpoints/2025-08-24_08-31-23/last_ckpt_infer.pt", max_sequences=40, tokenizer_path="s3://suno-data/georg/models/tokenizers/tokenizer_60k.json", compile=False, port=18865, ) # Get model config cfg = client.get_model_cfg() print(f"Model config loaded: {cfg}") # Test lyrics (simple, non-copyrighted content) lyrics = """ [verse] walking down the street today sunshine brightening up my way feeling good and feeling free this is how life's meant to be [chorus] every step brings something new every day a different view life is good when you believe in the dreams you can achieve """ # Create generation config gconf = GenerationConfig( text=lyrics, text_tags="", cfg_coef=2.0, cfg_coef_tags=3.0, cfg_coef_max_steps=100, cfg_coef_tags_max_steps=150, n_batch=1, min_text_offset=0, eos_pad_duration_s=0, max_gen_duration_s=max_gen_duration_s, ) # Create requests print(f"Creating {num_requests} requests...") import uuid unique_prefix = str(uuid.uuid4())[:8] requests = [ make_request(f"test_{unique_prefix}_{i}", gconf, cfg, client.get_tokenizer()) for i in range(num_requests) ] # Debug: check allow_eos setting and min generation duration print(f"First request allow_eos: {requests[0].allow_eos}") print(f"First request min_gen_duration_s: {requests[0].min_gen_duration_s}") print(f"First request max_gen_duration_s: {requests[0].max_gen_duration_s}") print( f"GenerationConfig min_gen_duration_s: {getattr(gconf, 'min_gen_duration_s', 'not set')}" ) print(f"Semantic pad token: {cfg.semantic_pad_token}") # Submit jobs jobs = [] print("Submitting jobs...") for request in tqdm(requests, desc="Adding requests"): job_id = client.add_request(request) jobs.append(job_id) print(f"Added job {job_id}") # Wait for completion print("Waiting for job completion...") start_time = time.time() while True: all_completed = True for job in jobs: state = client.get_job_state(job) if state is None: print(f"Warning: Job {job} state is None") continue if not state.get("completed", False): all_completed = False break if all_completed: print("All jobs completed!") break # Timeout after 5 minutes if time.time() - start_time > 300: print("Timeout waiting for jobs to complete") break time.sleep(1) # Check results print("\nJob results:") for i, job in enumerate(jobs): state = client.get_job_state(job) if state: completed = state["completed"] eos_step = state.get("eos_step") ttl = state.get("ttl", 0) print( f"Job {i}: completed={completed}, eos_step={eos_step}, ttl={ttl:.2f}ms" ) # Get generated codes codes = client.get_generated_codes(job) if codes is not None and hasattr(codes, "shape"): shape = codes.shape print(f" Generated codes shape: {shape}") # Handle both (T, 2) and (2, T) formats num_tokens = shape[1] if shape[0] == 2 else shape[0] if num_tokens > 100: print(f" ✓ Generated sufficient tokens ({num_tokens} tokens)") else: print( f" ⚠ Generated fewer tokens than expected ({num_tokens} tokens)" ) else: print(" ⚠ No codes generated or invalid format") else: print(f"Job {i}: No state available") # Cleanup print("\nCleaning up...") try: client.terminate_processes() print("✓ BCT engine processes terminated successfully") except Exception as e: print(f"⚠ Error during cleanup: {e}") print("BCT engine test completed!") # Wait for trace daemon to finish and export print("Waiting for trace daemon to export traces...") daemon_proc.join(timeout=130) # Wait up to 130 seconds for daemon to finish if daemon_proc.is_alive(): print("Trace daemon still running, terminating...") daemon_proc.terminate() daemon_proc.join() print("✓ Trace export completed! Check /tmp for trace files.") exit(0) if __name__ == "__main__": parser = argparse.ArgumentParser(description="Test BCT engine integration") parser.add_argument( "--num-requests", type=int, default=4, help="Number of requests to submit (default: 4)", ) parser.add_argument( "--max-gen-duration", type=float, default=30.0, help="Maximum generation duration in seconds (default: 30.0)", ) args = parser.parse_args() try: test_bct_engine( num_requests=args.num_requests, max_gen_duration_s=args.max_gen_duration ) except KeyboardInterrupt: print("\nTest interrupted by user") except Exception as e: print(f"Test failed with error: {e}") import traceback traceback.print_exc()