Team Ai
Apppublic

AUXteam/Critical_Code_Agent

sourceHugging Faceupdated 8mo agoView on Hugging Face
0likes
launch_scientist.py421 linesDownload Raw Back to root
1import argparse2import json3import multiprocessing4import openai5import os6import os.path as osp7import shutil8import sys9import time10import torch11from aider.coders import Coder12from aider.io import InputOutput13from aider.models import Model14from datetime import datetime15 16from ai_scientist.generate_ideas import generate_ideas, check_idea_novelty17from ai_scientist.llm import create_client, AVAILABLE_LLMS18from ai_scientist.perform_experiments import perform_experiments19from ai_scientist.perform_review import perform_review, load_paper, perform_improvement20from ai_scientist.perform_writeup import perform_writeup, generate_latex21 22NUM_REFLECTIONS = 323 24 25def print_time():26    print(datetime.now().strftime("%Y-%m-%d %H:%M:%S"))27 28 29def parse_arguments():30    parser = argparse.ArgumentParser(description="Run AI scientist experiments")31    parser.add_argument(32        "--skip-idea-generation",33        action="store_true",34        help="Skip idea generation and load existing ideas",35    )36    parser.add_argument(37        "--skip-novelty-check",38        action="store_true",39        help="Skip novelty check and use existing ideas",40    )41    # add type of experiment (nanoGPT, Boston, etc.)42    parser.add_argument(43        "--experiment",44        type=str,45        default="nanoGPT",46        help="Experiment to run AI Scientist on.",47    )48    parser.add_argument(49        "--model",50        type=str,51        default="claude-3-5-sonnet-20240620",52        choices=AVAILABLE_LLMS,53        help="Model to use for AI Scientist.",54    )55    parser.add_argument(56        "--writeup",57        type=str,58        default="latex",59        choices=["latex"],60        help="What format to use for writeup",61    )62    parser.add_argument(63        "--parallel",64        type=int,65        default=0,66        help="Number of parallel processes to run. 0 for sequential execution.",67    )68    parser.add_argument(69        "--improvement",70        action="store_true",71        help="Improve based on reviews.",72    )73    parser.add_argument(74        "--gpus",75        type=str,76        default=None,77        help="Comma-separated list of GPU IDs to use (e.g., '0,1,2'). If not specified, all available GPUs will be used.",78    )79    parser.add_argument(80        "--num-ideas",81        type=int,82        default=50,83        help="Number of ideas to generate",84    )85    parser.add_argument(86        "--engine",87        type=str,88        default="semanticscholar",89        choices=["semanticscholar", "openalex"],90        help="Scholar engine to use.",91    )92    return parser.parse_args()93 94 95def get_available_gpus(gpu_ids=None):96    if gpu_ids is not None:97        return [int(gpu_id) for gpu_id in gpu_ids.split(",")]98    return list(range(torch.cuda.device_count()))99 100 101def check_latex_dependencies():102    """103    Check if required LaTeX dependencies are installed on the system.104    Returns True if all dependencies are found, False otherwise.105    """106    import shutil107    import sys108 109    required_dependencies = ['pdflatex', 'chktex']110    missing_deps = []111 112    for dep in required_dependencies:113        if shutil.which(dep) is None:114            missing_deps.append(dep)115    116    if missing_deps:117        print("Error: Required LaTeX dependencies not found:", file=sys.stderr)118        return False119    120    return True121    122def worker(123        queue,124        base_dir,125        results_dir,126        model,127        client,128        client_model,129        writeup,130        improvement,131        gpu_id,132):133    os.environ["CUDA_VISIBLE_DEVICES"] = str(gpu_id)134    print(f"Worker {gpu_id} started.")135    while True:136        idea = queue.get()137        if idea is None:138            break139        success = do_idea(140            base_dir,141            results_dir,142            idea,143            model,144            client,145            client_model,146            writeup,147            improvement,148            log_file=True,149        )150        print(f"Completed idea: {idea['Name']}, Success: {success}")151    print(f"Worker {gpu_id} finished.")152 153 154def do_idea(155        base_dir,156        results_dir,157        idea,158        model,159        client,160        client_model,161        writeup,162        improvement,163        log_file=False,164):165    ## CREATE PROJECT FOLDER166    timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")167    idea_name = f"{timestamp}_{idea['Name']}"168    folder_name = osp.join(results_dir, idea_name)169    assert not osp.exists(folder_name), f"Folder {folder_name} already exists."170    destination_dir = folder_name171    shutil.copytree(base_dir, destination_dir, dirs_exist_ok=True)172    with open(osp.join(base_dir, "run_0", "final_info.json"), "r") as f:173        baseline_results = json.load(f)174    # Check if baseline_results is a dictionary before extracting means175    if isinstance(baseline_results, dict):176        baseline_results = {k: v["means"] for k, v in baseline_results.items()}177    exp_file = osp.join(folder_name, "experiment.py")178    vis_file = osp.join(folder_name, "plot.py")179    notes = osp.join(folder_name, "notes.txt")180    with open(notes, "w") as f:181        f.write(f"# Title: {idea['Title']}\n")182        f.write(f"# Experiment description: {idea['Experiment']}\n")183        f.write(f"## Run 0: Baseline\n")184        f.write(f"Results: {baseline_results}\n")185        f.write(f"Description: Baseline results.\n")186    if log_file:187        original_stdout = sys.stdout188        original_stderr = sys.stderr189        log_path = osp.join(folder_name, "log.txt")190        log = open(log_path, "a")191        sys.stdout = log192        sys.stderr = log193    try:194        print_time()195        print(f"*Starting idea: {idea_name}*")196        ## PERFORM EXPERIMENTS197        fnames = [exp_file, vis_file, notes]198        io = InputOutput(199            yes=True, chat_history_file=f"{folder_name}/{idea_name}_aider.txt"200        )201        if model == "deepseek-coder-v2-0724":202            main_model = Model("deepseek/deepseek-coder")203        elif model == "deepseek-reasoner":204            main_model = Model("deepseek/deepseek-reasoner")205        elif model == "llama3.1-405b":206            main_model = Model("openrouter/meta-llama/llama-3.1-405b-instruct")207        else:208            main_model = Model(model)209        coder = Coder.create(210            main_model=main_model,211            fnames=fnames,212            io=io,213            stream=False,214            use_git=False,215            edit_format="diff",216        )217 218        print_time()219        print(f"*Starting Experiments*")220        try:221            success = perform_experiments(idea, folder_name, coder, baseline_results)222        except Exception as e:223            print(f"Error during experiments: {e}")224            print(f"Experiments failed for idea {idea_name}")225            return False226 227        if not success:228            print(f"Experiments failed for idea {idea_name}")229            return False230 231        print_time()232        print(f"*Starting Writeup*")233        ## PERFORM WRITEUP234        if writeup == "latex":235            writeup_file = osp.join(folder_name, "latex", "template.tex")236            fnames = [exp_file, writeup_file, notes]237            if model == "deepseek-coder-v2-0724":238                main_model = Model("deepseek/deepseek-coder")239            elif model == "deepseek-reasoner":240                main_model = Model("deepseek/deepseek-reasoner")241            elif model == "llama3.1-405b":242                main_model = Model("openrouter/meta-llama/llama-3.1-405b-instruct")243            else:244                main_model = Model(model)245            coder = Coder.create(246                main_model=main_model,247                fnames=fnames,248                io=io,249                stream=False,250                use_git=False,251                edit_format="diff",252            )253            try:254                perform_writeup(idea, folder_name, coder, client, client_model, engine=args.engine)255            except Exception as e:256                print(f"Failed to perform writeup: {e}")257                return False258            print("Done writeup")259        else:260            raise ValueError(f"Writeup format {writeup} not supported.")261 262        print_time()263        print(f"*Starting Review*")264        ## REVIEW PAPER265        if writeup == "latex":266            try:267                paper_text = load_paper(f"{folder_name}/{idea['Name']}.pdf")268                review = perform_review(269                    paper_text,270                    model="gpt-4o-2024-05-13",271                    client=openai.OpenAI(),272                    num_reflections=5,273                    num_fs_examples=1,274                    num_reviews_ensemble=5,275                    temperature=0.1,276                )277                # Store the review in separate review.txt file278                with open(osp.join(folder_name, "review.txt"), "w") as f:279                    f.write(json.dumps(review, indent=4))280            except Exception as e:281                print(f"Failed to perform review: {e}")282                return False283 284        ## IMPROVE WRITEUP285        if writeup == "latex" and improvement:286            print_time()287            print(f"*Starting Improvement*")288            try:289                perform_improvement(review, coder)290                generate_latex(291                    coder, folder_name, f"{folder_name}/{idea['Name']}_improved.pdf"292                )293                paper_text = load_paper(f"{folder_name}/{idea['Name']}_improved.pdf")294                review = perform_review(295                    paper_text,296                    model="gpt-4o-2024-05-13",297                    client=openai.OpenAI(),298                    num_reflections=5,299                    num_fs_examples=1,300                    num_reviews_ensemble=5,301                    temperature=0.1,302                )303                # Store the review in separate review.txt file304                with open(osp.join(folder_name, "review_improved.txt"), "w") as f:305                    f.write(json.dumps(review))306            except Exception as e:307                print(f"Failed to perform improvement: {e}")308                return False309        return True310    except Exception as e:311        print(f"Failed to evaluate idea {idea_name}: {str(e)}")312        return False313    finally:314        print("FINISHED IDEA")315        if log_file:316            sys.stdout = original_stdout317            sys.stderr = original_stderr318            log.close()319 320 321if __name__ == "__main__":322    args = parse_arguments()323 324    # Check available GPUs and adjust parallel processes if necessary325    available_gpus = get_available_gpus(args.gpus)326    if args.parallel > len(available_gpus):327        print(328            f"Warning: Requested {args.parallel} parallel processes, but only {len(available_gpus)} GPUs available. Adjusting to {len(available_gpus)}."329        )330        args.parallel = len(available_gpus)331 332    print(f"Using GPUs: {available_gpus}")333 334    # Check LaTeX dependencies before proceeding335    if args.writeup == "latex" and not check_latex_dependencies():336        sys.exit(1)337 338    # Create client339    client, client_model = create_client(args.model)340 341    base_dir = osp.join("templates", args.experiment)342    results_dir = osp.join("results", args.experiment)343    ideas = generate_ideas(344        base_dir,345        client=client,346        model=client_model,347        skip_generation=args.skip_idea_generation,348        max_num_generations=args.num_ideas,349        num_reflections=NUM_REFLECTIONS,350    )351    if not args.skip_novelty_check:352        ideas = check_idea_novelty(353            ideas,354            base_dir=base_dir,355            client=client,356            model=client_model,357            engine=args.engine,358        )359 360    with open(osp.join(base_dir, "ideas.json"), "w") as f:361        json.dump(ideas, f, indent=4)362 363    novel_ideas = [idea for idea in ideas if idea["novel"]]364    # novel_ideas = list(reversed(novel_ideas))365 366    if args.parallel > 0:367        print(f"Running {args.parallel} parallel processes")368        queue = multiprocessing.Queue()369        for idea in novel_ideas:370            queue.put(idea)371 372        processes = []373        for i in range(args.parallel):374            gpu_id = available_gpus[i % len(available_gpus)]375            p = multiprocessing.Process(376                target=worker,377                args=(378                    queue,379                    base_dir,380                    results_dir,381                    args.model,382                    client,383                    client_model,384                    args.writeup,385                    args.improvement,386                    gpu_id,387                ),388            )389            p.start()390            time.sleep(150)391            processes.append(p)392 393        # Signal workers to exit394        for _ in range(args.parallel):395            queue.put(None)396 397        for p in processes:398            p.join()399 400        print("All parallel processes completed.")401    else:402        for idea in novel_ideas:403            print(f"Processing idea: {idea['Name']}")404            try:405                success = do_idea(406                    base_dir,407                    results_dir,408                    idea,409                    args.model,410                    client,411                    client_model,412                    args.writeup,413                    args.improvement,414                )415                print(f"Completed idea: {idea['Name']}, Success: {success}")416            except Exception as e:417                print(f"Failed to evaluate idea {idea['Name']}: {str(e)}")418                import traceback419                print(traceback.format_exc())420    print("All ideas evaluated.")421