11import json
2+ import math
23import os
34import pickle
45import re
56import shutil
6- import tarfile
77import time
88from pathlib import Path
99from typing import Any , Dict , List , Optional , Tuple
2626from rdagent .scenarios .data_science .experiment .experiment import DSExperiment
2727from rdagent .utils .agent .ret import PythonAgentOut
2828from rdagent .utils .agent .tpl import T
29+ from rdagent .utils .archive import safe_extract_tar
2930from rdagent .utils .fmt import shrink_text
3031from rdagent .utils .workflow import wait_retry
3132
@@ -350,6 +351,34 @@ def print_code(self, data_py_code: str, grade_py_code: str):
350351 print (grade_py_code )
351352 print ("======== code end ========" )
352353
354+ def _validate_grade_script (self , grade_py_code : str , reference_exp : DSExperiment , mock_folder : str ) -> None :
355+ """Run a reusable grade script and verify its output contract."""
356+ input_folder = T ("scenarios.data_science.share:scen.input_path" ).r ()
357+ submission_path = Path (mock_folder ) / "submission.csv"
358+ if not submission_path .exists ():
359+ message = f"Cannot validate grade.py because { submission_path } does not exist."
360+ raise RuntimeError (message )
361+
362+ ws = FBWorkspace ()
363+ ws .inject_code_from_file_dict (reference_exp .experiment_workspace )
364+ ws .inject_files (** {"grade.py" : grade_py_code })
365+ shutil .copy (str (submission_path ), str (ws .workspace_path / "submission.csv" ))
366+ env = get_ds_env (extra_volumes = {str (Path (mock_folder ) / input_folder ): {"bind" : input_folder , "mode" : "rw" }})
367+ result = ws .run (env = env , entry = f"python grade.py --cache-buster={ time .time ()} " )
368+ stdout = re .sub (r"^chmod:.*\n?" , "" , result .stdout , flags = re .MULTILINE )
369+
370+ if result .exit_code != 0 :
371+ output = shrink_text (stdout , context_lines = 20 , line_len = 500 )
372+ message = f"grade.py validation failed with exit code { result .exit_code } : { output } "
373+ raise RuntimeError (message )
374+ if _parsing_score (stdout ) is None :
375+ output = shrink_text (stdout , context_lines = 20 , line_len = 500 )
376+ message = (
377+ "grade.py must print a valid JSON object whose 'score' is a finite numeric value; "
378+ f"received stdout: { output } "
379+ )
380+ raise RuntimeError (message )
381+
353382 def _prepare_validation_scripts (
354383 self , reference_exp : DSExperiment , competition : str , mock_folder : str
355384 ) -> Tuple [str , str ]:
@@ -361,6 +390,7 @@ def _prepare_validation_scripts(
361390 data_py_path = Path (mock_folder ) / "data.py"
362391 grade_py_path = Path (mock_folder ) / "grade.py"
363392 label_path = Path (mock_folder ) / "workspace_input/label.csv"
393+ submission_path = Path (mock_folder ) / "submission.csv"
364394 reference_code = reference_exp .experiment_workspace .file_dict .get ("main.py" , "" )
365395 if not reference_code :
366396 raise RuntimeError ("ValidationSelector: No code found in the reference experiment." )
@@ -370,7 +400,7 @@ def _prepare_validation_scripts(
370400 shutil .copy (self .sample_code_path / competition / "grade.py" , grade_py_path )
371401 data_py_code = data_py_path .read_text ()
372402 grade_py_code = grade_py_path .read_text ()
373- if not label_path .exists ():
403+ if not label_path .exists () or not submission_path . exists () :
374404 ws = FBWorkspace ()
375405 if self .sample_rate != 0.8 :
376406 data_py_code = data_py_code .replace ("0.8" , str (self .sample_rate )).replace (
@@ -390,10 +420,11 @@ def _prepare_validation_scripts(
390420 ) # Do not cache the result
391421 if result .exit_code == 0 :
392422 self .print_code (data_py_code , grade_py_code )
423+ self ._validate_grade_script (grade_py_code , reference_exp , mock_folder )
393424 return data_py_code , grade_py_code
394425
395426 # --- Generate data.py if needed ---
396- if not data_py_path .exists () or not label_path .exists ():
427+ if not data_py_path .exists () or not label_path .exists () or not submission_path . exists () :
397428 logger .info (f"Generating synthetic data script: { data_py_path } " )
398429 data_py_code = self ._generate_and_run_script (
399430 script_type = "data" ,
@@ -408,23 +439,31 @@ def _prepare_validation_scripts(
408439 data_py_code = data_py_path .read_text ()
409440
410441 # --- Generate grade.py if needed ---
411- if not grade_py_path .exists ():
412- logger .info (f"Generating grading script: { grade_py_path } " )
413- grade_py_code = self ._generate_and_run_script (
414- script_type = "grade" ,
415- prompt_template_key = "grade" ,
416- reference_exp = reference_exp ,
417- competition = competition ,
418- mock_folder = mock_folder ,
419- prompt_kwargs = {
420- "reference_code" : reference_code ,
421- "sample_code" : data_py_code ,
422- "input_folder" : input_folder ,
423- },
424- )
425- grade_py_path .write_text (grade_py_code )
426- self .print_code (data_py_code , grade_py_code )
427- return data_py_code , grade_py_path .read_text ()
442+ if grade_py_path .exists ():
443+ grade_py_code = grade_py_path .read_text ()
444+ try :
445+ self ._validate_grade_script (grade_py_code , reference_exp , mock_folder )
446+ except RuntimeError as exc :
447+ logger .warning (f"Cached grade.py is incompatible and will be regenerated: { exc } " )
448+ else :
449+ return data_py_code , grade_py_code
450+
451+ logger .info (f"Generating grading script: { grade_py_path } " )
452+ grade_py_code = self ._generate_and_run_script (
453+ script_type = "grade" ,
454+ prompt_template_key = "grade" ,
455+ reference_exp = reference_exp ,
456+ competition = competition ,
457+ mock_folder = mock_folder ,
458+ prompt_kwargs = {
459+ "reference_code" : reference_code ,
460+ "sample_code" : data_py_code ,
461+ "input_folder" : input_folder ,
462+ },
463+ )
464+ grade_py_path .write_text (grade_py_code )
465+ self .print_code (data_py_code , grade_py_code )
466+ return data_py_code , grade_py_code
428467
429468 def _generate_and_run_script (
430469 self ,
@@ -582,19 +621,14 @@ def _parsing_score(grade_stdout: str) -> Optional[float]:
582621 continue
583622 json_str = m .group (0 )
584623 try :
585- # Priority 1: JSON parsing
586- return float (json .loads (json_str )["score" ])
587- except :
588- pass
589- try :
590- # Priority 2: Eval dict
591- return float (eval (json_str )["score" ])
592- except :
593- pass
594- try :
595- # Priority 3: Regex for the last number in the string
596- return float (re .findall (r"[-+]?\d*\.\d+|\d+" , json_str )[- 1 ])
597- except :
624+ score = json .loads (json_str )["score" ]
625+ if isinstance (score , bool ) or not isinstance (score , (int , float )):
626+ continue
627+ score = float (score )
628+ if not math .isfinite (score ):
629+ continue
630+ return score
631+ except (KeyError , TypeError , ValueError ):
598632 pass
599633 return None
600634
@@ -626,9 +660,8 @@ def try_get_loop_id(trace: Trace, exp: DSExperiment):
626660 return index
627661
628662
629- def extract_tar (tar_path : str , to_dir : str = "log" ) -> str :
630- with tarfile .open (tar_path , mode = "r:*" ) as tar :
631- tar .extractall (path = to_dir )
663+ def extract_tar (tar_path : str , to_dir : str = "log" ) -> None :
664+ safe_extract_tar (tar_path , to_dir )
632665
633666
634667# ==============================================================================
0 commit comments