Akki0404 commited on
Commit
25f7d76
·
1 Parent(s): 9dca18f

fix log_end missing score field

Browse files
Files changed (1) hide show
  1. inference.py +87 -1
inference.py CHANGED
@@ -304,6 +304,88 @@ async def run_task(client: OpenAI, task_name: str):
304
  score=score,
305
  rewards=rewards,
306
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
307
  async def main():
308
  client = OpenAI(base_url=API_BASE_URL, api_key=HF_TOKEN)
309
  tasks = [
@@ -312,9 +394,13 @@ async def main():
312
  "adversarial_detection",
313
  "streaming_detection",
314
  "phonecall_detection",
 
315
  ]
316
  for task in tasks:
317
- await run_task(client, task)
 
 
 
318
 
319
 
320
  if __name__ == "__main__":
 
304
  score=score,
305
  rewards=rewards,
306
  )
307
+
308
+
309
+ async def run_realtime_task(client: OpenAI, task_name: str):
310
+ """Run one episode of realtime_detection.
311
+
312
+ Strategy: gather 2 features (temporal + spectral) then classify
313
+ immediately to minimize the time penalty (-0.03 per extra step).
314
+ The agent only takes 3 steps total: 2 gathering + 1 classify.
315
+ """
316
+ rewards: List[float] = []
317
+ steps_taken = 0
318
+ success = False
319
+ score = 0.05
320
+ context = {}
321
+
322
+ log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)
323
+
324
+ try:
325
+ # Reset
326
+ reset_response = env_reset(task_name)
327
+ observation = reset_response.get("observation", {})
328
+ context = {
329
+ "task_name": observation.get("task_name", task_name),
330
+ "difficulty": observation.get("difficulty", ""),
331
+ "visible_features": {},
332
+ "comparison_result": None,
333
+ "evidence_summary": None,
334
+ "actions_taken": [],
335
+ }
336
+
337
+ # Step 1: Request temporal features
338
+ action1 = {"action_type": "request_temporal_features"}
339
+ step1 = env_step(action1, task_name)
340
+ observation = step1.get("observation", {})
341
+ reward1 = _clamp_score(float(step1.get("reward", 0.05)))
342
+ rewards.append(reward1)
343
+ steps_taken = 1
344
+ context["visible_features"] = observation.get("visible_features", {})
345
+ context["actions_taken"] = observation.get("actions_taken", [])
346
+ log_step(step=1, action=action1, reward=reward1,
347
+ done=step1.get("done", False), error=None)
348
+
349
+ # Step 2: Request spectral features
350
+ action2 = {"action_type": "request_spectral_features"}
351
+ step2 = env_step(action2, task_name)
352
+ observation = step2.get("observation", {})
353
+ reward2 = _clamp_score(float(step2.get("reward", 0.05)))
354
+ rewards.append(reward2)
355
+ steps_taken = 2
356
+ context["visible_features"] = observation.get("visible_features", {})
357
+ context["actions_taken"] = observation.get("actions_taken", [])
358
+ log_step(step=2, action=action2, reward=reward2,
359
+ done=step2.get("done", False), error=None)
360
+
361
+ # Step 3: Classify immediately (no extra steps = no time penalty)
362
+ classification = get_classification(client, context)
363
+ action3 = {
364
+ "action_type": "final_classify",
365
+ "label": classification["label"],
366
+ "confidence": classification["confidence"],
367
+ "reasoning": classification.get("reasoning", ""),
368
+ }
369
+ step3 = env_step(action3, task_name)
370
+ reward3 = _clamp_score(float(step3.get("reward", 0.05)))
371
+ rewards.append(reward3)
372
+ steps_taken = 3
373
+ log_step(step=3, action=action3, reward=reward3,
374
+ done=step3.get("done", True), error=None)
375
+
376
+ score = reward3
377
+ success = score >= SUCCESS_SCORE_THRESHOLD
378
+
379
+ except Exception as e:
380
+ print(f"[DEBUG] Task error: {e}", flush=True)
381
+
382
+ finally:
383
+ log_end(
384
+ success=success,
385
+ steps=steps_taken,
386
+ score=score,
387
+ rewards=rewards,
388
+ )
389
  async def main():
390
  client = OpenAI(base_url=API_BASE_URL, api_key=HF_TOKEN)
391
  tasks = [
 
394
  "adversarial_detection",
395
  "streaming_detection",
396
  "phonecall_detection",
397
+ "realtime_detection",
398
  ]
399
  for task in tasks:
400
+ if task == "realtime_detection":
401
+ await run_realtime_task(client, task)
402
+ else:
403
+ await run_task(client, task)
404
 
405
 
406
  if __name__ == "__main__":