run_full_eval.py 7.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219
  1. import sys
  2. import os
  3. import time
  4. import requests
  5. import argparse
  6. from datetime import datetime
  7. # Ensure project root is in python path
  8. sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
  9. from sqlalchemy.orm import Session
  10. from app.db.session import SessionLocal
  11. from app.crud import create_forum, get_forum
  12. from app.schemas import ForumCreate
  13. # Import our new tools
  14. # Ensure project root is in python path
  15. project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
  16. if project_root not in sys.path:
  17. sys.path.append(project_root)
  18. # And also add exam folder itself
  19. exam_dir = os.path.dirname(os.path.abspath(__file__))
  20. if exam_dir not in sys.path:
  21. sys.path.append(exam_dir)
  22. # Import assuming running from project root or inside exam/
  23. try:
  24. from exam.generate_roles import generate_roles
  25. from exam.baseline_eval import create_baseline_forum
  26. from exam.standard_eval import evaluate_forum
  27. from exam.ablation_study import compare_forums
  28. except ImportError:
  29. # Fallback for direct execution
  30. from generate_roles import generate_roles
  31. from baseline_eval import create_baseline_forum
  32. from standard_eval import evaluate_forum
  33. from ablation_study import compare_forums
  34. # Configuration
  35. API_BASE_URL = "http://localhost:8000/api/v1"
  36. USERNAME = "experiment_admin"
  37. PASSWORD = "admin_password"
  38. def get_token():
  39. # Simple login as experiment_admin or create
  40. login_url = f"{API_BASE_URL}/auth/login"
  41. payload = {"username": USERNAME, "password": PASSWORD}
  42. try:
  43. resp = requests.post(login_url, data=payload)
  44. if resp.status_code == 200:
  45. return resp.json()["access_token"]
  46. # Try registering
  47. reg_url = f"{API_BASE_URL}/auth/register"
  48. requests.post(reg_url, json=payload)
  49. resp = requests.post(login_url, data=payload)
  50. if resp.status_code == 200:
  51. return resp.json()["access_token"]
  52. print(f"Failed to login/register as {USERNAME}.")
  53. return None
  54. except:
  55. print("Backend not running?")
  56. return None
  57. def run_standard_forum(topic: str, persona_ids: list, duration_minutes: int = 5, ablation_flags: dict = None):
  58. """
  59. Creates a forum, starts it via API, and waits for completion.
  60. """
  61. db = SessionLocal()
  62. try:
  63. if ablation_flags:
  64. print(f"Creating Forum with Ablation Flags: {ablation_flags}...")
  65. else:
  66. print(f"Creating Standard Forum: '{topic}' with {len(persona_ids)} agents...")
  67. # We need the user ID for creator_id.
  68. from app.crud import get_user_by_username
  69. user = get_user_by_username(db, USERNAME)
  70. if not user:
  71. print(f"User {USERNAME} not found in DB.")
  72. return None
  73. # 1. Create Forum in DB
  74. f_create = ForumCreate(
  75. topic=topic,
  76. participant_ids=persona_ids,
  77. duration_minutes=duration_minutes,
  78. moderator_id=persona_ids[0] if persona_ids else 1 # Default to first agent
  79. )
  80. forum = create_forum(db, f_create, user.id)
  81. print(f"Forum Created (ID: {forum.id}). Duration: {duration_minutes} min.")
  82. # 2. Start Forum via API
  83. token = get_token()
  84. if not token:
  85. print("Cannot get API token. Is backend running?")
  86. return None
  87. start_url = f"{API_BASE_URL}/forums/{forum.id}/start"
  88. headers = {"Authorization": f"Bearer {token}"}
  89. # Pass ablation flags
  90. payload = {}
  91. if ablation_flags:
  92. payload["ablation_flags"] = ablation_flags
  93. resp = requests.post(start_url, json=payload, headers=headers)
  94. if resp.status_code != 200:
  95. print(f"Failed to start forum: {resp.text}")
  96. return None
  97. print("Forum started. Waiting for completion...")
  98. # 3. Wait for completion
  99. # Poll DB status
  100. while True:
  101. db.expire_all() # Refresh
  102. f = get_forum(db, forum.id)
  103. if not f:
  104. print("Forum disappeared?")
  105. break
  106. status = f.status
  107. print(f" Status: {status} (Time: {datetime.now().strftime('%H:%M:%S')})")
  108. if status == "completed":
  109. print("Forum Completed!")
  110. break
  111. elif status == "closed":
  112. print("Forum Closed (Time's up)!")
  113. break
  114. elif status == "failed":
  115. print("Forum Failed!")
  116. break
  117. time.sleep(10) # Poll every 10s
  118. return forum.id
  119. finally:
  120. db.close()
  121. def run_full_evaluation(topic: str, num_agents: int = 3, duration: int = 5):
  122. print("="*50)
  123. print(f"STARTING FULL EVALUATION PIPELINE")
  124. print(f"Topic: {topic}")
  125. print("="*50)
  126. # Step 1: Generate Roles
  127. print("\n[Step 1] Generating Roles...")
  128. persona_ids = generate_roles(topic, n=num_agents, owner_username=USERNAME)
  129. if not persona_ids:
  130. print("Failed to generate roles.")
  131. return
  132. # Step 2: Run Standard Forum
  133. print("\n[Step 2] Running Standard Multi-Agent Forum...")
  134. std_forum_id = run_standard_forum(topic, persona_ids, duration_minutes=duration)
  135. if not std_forum_id:
  136. print("Failed to run standard forum.")
  137. return
  138. # Step 3: Generate Baseline
  139. print("\n[Step 3] Generating Single LLM Baseline...")
  140. baseline_forum_id = create_baseline_forum(topic, owner_username=USERNAME)
  141. if not baseline_forum_id:
  142. print("Failed to generate baseline.")
  143. return
  144. # Step 4: Run Ablation Forums
  145. print("\n[Step 4.1] Running Ablation: No Summary...")
  146. no_summary_id = run_standard_forum(topic, persona_ids, duration_minutes=duration, ablation_flags={"no_summary": True})
  147. print("\n[Step 4.2] Running Ablation: No Private Memory...")
  148. no_private_id = run_standard_forum(topic, persona_ids, duration_minutes=duration, ablation_flags={"no_private_memory": True})
  149. print("\n[Step 4.3] Running Ablation: No Shared Memory...")
  150. no_shared_id = run_standard_forum(topic, persona_ids, duration_minutes=duration, ablation_flags={"no_shared_memory": True})
  151. # Step 5: Evaluations
  152. print("\n[Step 5] Running Comparisons...")
  153. # 5.1 Standard vs Baseline (Original request)
  154. print("\n>>> Standard vs Baseline")
  155. compare_forums(std_forum_id, baseline_forum_id, "Multi-Agent Discussion vs Single LLM Baseline")
  156. # 5.2 Standard vs No Summary
  157. if no_summary_id:
  158. print("\n>>> Standard vs No Summary")
  159. compare_forums(std_forum_id, no_summary_id, "Standard vs No Periodic Summary")
  160. # 5.3 Standard vs No Private Memory
  161. if no_private_id:
  162. print("\n>>> Standard vs No Private Memory")
  163. compare_forums(std_forum_id, no_private_id, "Standard vs No Private Memory (Stateless Agents)")
  164. # 5.4 Standard vs No Shared Memory
  165. if no_shared_id:
  166. print("\n>>> Standard vs No Shared Memory")
  167. compare_forums(std_forum_id, no_shared_id, "Standard vs No Shared Memory Context")
  168. print("\n" + "="*50)
  169. print("EVALUATION PIPELINE COMPLETE")
  170. print("Check exam/results/ for reports.")
  171. print("="*50)
  172. if __name__ == "__main__":
  173. parser = argparse.ArgumentParser(description="Run full evaluation pipeline.")
  174. parser.add_argument("--topic", type=str, default="人工智能是否应该拥有人权?", help="Topic for evaluation.")
  175. parser.add_argument("--agents", type=int, default=3, help="Number of agents.")
  176. parser.add_argument("--duration", type=int, default=5, help="Duration in minutes.")
  177. args = parser.parse_args()
  178. run_full_evaluation(args.topic, args.agents, args.duration)