From 77755533eb155350ab7257b25305c0a1158d4294 Mon Sep 17 00:00:00 2001 From: MagellaX Date: Thu, 14 Aug 2025 18:03:03 +0530 Subject: [PATCH 1/4] feat(gepa): integrate GEPA via rl --gepa; add acceptance gate, pareto selection, env prompt overrides; add dry-run for CPU --- configs/reverse_text/gepa.toml | 34 +++ .../vf_hendrycks_math/vf_hendrycks_math.py | 8 +- .../vf_reverse_text/vf_reverse_text.py | 5 +- pyproject.toml | 1 + src/prime_rl/optimizer/__init__.py | 8 + src/prime_rl/optimizer/gepa/config.py | 79 ++++++ src/prime_rl/optimizer/gepa/evaluate.py | 191 +++++++++++++ src/prime_rl/optimizer/gepa/gepa.py | 255 ++++++++++++++++++ src/prime_rl/optimizer/gepa/operators.py | 57 ++++ src/prime_rl/optimizer/gepa/reflection.py | 32 +++ src/prime_rl/optimizer/gepa/selection.py | 39 +++ src/prime_rl/rl.py | 13 + 12 files changed, 718 insertions(+), 4 deletions(-) create mode 100644 configs/reverse_text/gepa.toml create mode 100644 src/prime_rl/optimizer/__init__.py create mode 100644 src/prime_rl/optimizer/gepa/config.py create mode 100644 src/prime_rl/optimizer/gepa/evaluate.py create mode 100644 src/prime_rl/optimizer/gepa/gepa.py create mode 100644 src/prime_rl/optimizer/gepa/operators.py create mode 100644 src/prime_rl/optimizer/gepa/reflection.py create mode 100644 src/prime_rl/optimizer/gepa/selection.py diff --git a/configs/reverse_text/gepa.toml b/configs/reverse_text/gepa.toml new file mode 100644 index 0000000000..9cef32372f --- /dev/null +++ b/configs/reverse_text/gepa.toml @@ -0,0 +1,34 @@ +[model] +name = "Qwen/Qwen3-0.6B" + +[evaluate] +benchmarks = ["math500"] +subset_size = 32 +rollouts_per_prompt = 1 +max_tokens = 64 +min_tokens = 0 + +[operators] +mutation_rate = 0.7 +crossover_rate = 0.3 +max_prompt_chars = 2000 +enforce_diversity = true +min_levenshtein_distance = 40 + +[selection] +strategy = "top-k" +k = 6 +keep_elite = 2 + +population_size = 8 +generations = 3 +seed = 42 + +[log] +level = "info" + +[monitor] +# Enable W&B by adding a [monitor.wandb] section externally if desired +system_log_frequency = 0 + + diff --git a/environments/vf_hendrycks_math/vf_hendrycks_math.py b/environments/vf_hendrycks_math/vf_hendrycks_math.py index 8507d21aba..01e84b6fbf 100644 --- a/environments/vf_hendrycks_math/vf_hendrycks_math.py +++ b/environments/vf_hendrycks_math/vf_hendrycks_math.py @@ -2,7 +2,7 @@ from datasets import load_dataset -def load_environment(**kwargs) -> vf.Environment: +def load_environment(system_prompt: str | None = None, **kwargs) -> vf.Environment: import json from verifiers.utils.data_utils import extract_boxed_answer @@ -31,5 +31,9 @@ def correct_answer_reward_func(completion, info) -> float: weights=[1.0], ) - vf_env = vf.SingleTurnEnv(dataset=train_dataset, parser=parser, rubric=rubric) + # Pass optional system prompt if provided (SingleTurnEnv supports it) + if system_prompt is not None: + vf_env = vf.SingleTurnEnv(dataset=train_dataset, parser=parser, rubric=rubric, system_prompt=system_prompt) + else: + vf_env = vf.SingleTurnEnv(dataset=train_dataset, parser=parser, rubric=rubric) return vf_env diff --git a/environments/vf_reverse_text/vf_reverse_text.py b/environments/vf_reverse_text/vf_reverse_text.py index 8752f1b894..17b447e468 100644 --- a/environments/vf_reverse_text/vf_reverse_text.py +++ b/environments/vf_reverse_text/vf_reverse_text.py @@ -2,7 +2,7 @@ from datasets import load_dataset -def load_environment() -> vf.Environment: +def load_environment(system_prompt: str | None = None) -> vf.Environment: train_dataset = load_dataset("PrimeIntellect/Reverse-Text-RL", split="train").map( lambda x: { "question": x["prompt"], @@ -38,7 +38,8 @@ def lcs_ratio(x: str, y: str) -> float: weights=[1.0], ) - system_prompt = "Reverse the text character-by-character. Put your answer in tags." + base_system_prompt = "Reverse the text character-by-character. Put your answer in tags." + system_prompt = system_prompt or base_system_prompt vf_env = vf.SingleTurnEnv( dataset=train_dataset, diff --git a/pyproject.toml b/pyproject.toml index 60de0a22c8..a4027cf343 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -42,6 +42,7 @@ orchestrator = "prime_rl.orchestrator.orchestrator:main" inference = "prime_rl.inference.server:main" sft = "prime_rl.trainer.sft.train:main" eval = "prime_rl.eval.eval:main" +gepa = "prime_rl.optimizer.gepa.gepa:main" [build-system] requires = ["hatchling"] diff --git a/src/prime_rl/optimizer/__init__.py b/src/prime_rl/optimizer/__init__.py new file mode 100644 index 0000000000..285fcb0c0b --- /dev/null +++ b/src/prime_rl/optimizer/__init__.py @@ -0,0 +1,8 @@ +""" +Optimization utilities and algorithms (e.g., GEPA) for prompt and policy search. + +This namespace houses optimizers that orchestrate search procedures using the +existing inference, evaluation, and monitoring infrastructures. +""" + + diff --git a/src/prime_rl/optimizer/gepa/config.py b/src/prime_rl/optimizer/gepa/config.py new file mode 100644 index 0000000000..ce335e91aa --- /dev/null +++ b/src/prime_rl/optimizer/gepa/config.py @@ -0,0 +1,79 @@ +from pathlib import Path +from typing import Annotated, Literal + +from pydantic import Field + +from prime_rl.utils.config import LogConfig, MultiMonitorConfig, ModelConfig +from prime_rl.inference.config import InferenceConfig +from prime_rl.orchestrator.config import ClientConfig +from prime_rl.utils.pydantic_config import BaseConfig, BaseSettings + + +class EvaluateConfig(BaseConfig): + """Evaluation settings for GEPA prompt scoring.""" + + benchmark: Annotated[str, Field(description="Benchmark/dataset to evaluate on.")] = "math500" + rollouts_per_prompt: Annotated[int, Field(ge=1)] = 1 + pareto_size: Annotated[int, Field(ge=1, description="Number of instances in D_pareto.")] = 64 + feedback_pool_size: Annotated[int, Field(ge=1, description="Number of instances in D_feedback pool.")] = 256 + max_tokens: Annotated[int | None, Field(description="Max generated tokens per completion.")] = None + min_tokens: Annotated[int, Field(ge=0)] = 0 + + +class OperatorsConfig(BaseConfig): + """Mutation/crossover operator rates and constraints.""" + + mutation_rate: Annotated[float, Field(ge=0, le=1)] = 0.7 + crossover_rate: Annotated[float, Field(ge=0, le=1)] = 0.3 + max_prompt_chars: Annotated[int, Field(ge=1)] = 4000 + enforce_diversity: Annotated[bool, Field(description="Avoid near-duplicate prompts via distance checks.")] = True + min_levenshtein_distance: Annotated[int, Field(ge=0)] = 40 + + +class SelectionConfig(BaseConfig): + """Selection policy settings.""" + + strategy: Annotated[Literal["top-k", "tournament"], Field()] = "top-k" + k: Annotated[int, Field(ge=1)] = 8 + tournament_size: Annotated[int, Field(ge=1)] = 3 + keep_elite: Annotated[int, Field(ge=0)] = 2 + + +class GEPAConfig(BaseSettings): + """Top-level configuration for the GEPA optimizer.""" + + # Model/Client + model: ModelConfig = ModelConfig() + + # Evaluation + evaluate: EvaluateConfig = EvaluateConfig() + + # Evolution + population_size: Annotated[int, Field(ge=2)] = 16 + generations: Annotated[int, Field(ge=1)] = 10 + seed: Annotated[int | None, Field(description="Random seed for reproducibility.")] = 42 + + # Operators & selection + operators: OperatorsConfig = OperatorsConfig() + selection: SelectionConfig = SelectionConfig() + + # Budget & minibatch + budget_rollouts: Annotated[int, Field(ge=1, description="Total rollout budget for evolution.")] = 200 + minibatch_size: Annotated[int, Field(ge=1, description="Minibatch size b for acceptance test.")] = 8 + + # Logging/monitoring + log: LogConfig = LogConfig() + monitor: MultiMonitorConfig = MultiMonitorConfig() + + # IO + outputs_dir: Annotated[Path, Field(description="Directory for GEPA outputs.")] = Path("outputs") + run_name: Annotated[str | None, Field(description="Optional run name for outputs/W&B.")] = None + + # Dev + dry_run: Annotated[bool, Field(description="If True, skip real inference and fabricate scores.")] = False + + # Inference server (optional auto-spawn) and client + inference: InferenceConfig | None = None + client: ClientConfig = ClientConfig() + + diff --git a/src/prime_rl/optimizer/gepa/evaluate.py b/src/prime_rl/optimizer/gepa/evaluate.py new file mode 100644 index 0000000000..559458209b --- /dev/null +++ b/src/prime_rl/optimizer/gepa/evaluate.py @@ -0,0 +1,191 @@ +from __future__ import annotations + +import asyncio +from dataclasses import dataclass +from typing import Any + +from openai import AsyncOpenAI + +from prime_rl.eval.registry import get_benchmark_dataset +from prime_rl.orchestrator.client import generate_completion, setup_client +from prime_rl.orchestrator.config import ClientConfig, ModelConfig as OrchestratorModelConfig, SamplingConfig +from prime_rl.orchestrator.utils import compute_rewards, parse_completion_tokens, parse_completions +from prime_rl.utils.logger import get_logger + + +@dataclass +class PromptScore: + prompt: str + avg_reward: float + pass_at_k: float | None + avg_completion_len: float + meta: dict[str, Any] + + +async def _score_on_benchmark( + client: AsyncOpenAI, + benchmark: str, + system_prompt: str, + model_config: OrchestratorModelConfig, + sampling: SamplingConfig, + subset_size: int, + rollouts_per_prompt: int, +) -> tuple[float, float | None, float]: + logger = get_logger() + dataset = get_benchmark_dataset(benchmark) + dataset = dataset.select(range(min(len(dataset), subset_size))) + + prompts = [item["prompt"] for item in dataset] + prompts = [p for p in prompts for _ in range(rollouts_per_prompt)] + problem_ids = list(range(len(dataset))) + problem_ids = [pid for pid in problem_ids for _ in range(rollouts_per_prompt)] + + batch_messages = [[{"role": "system", "content": system_prompt}, {"role": "user", "content": p}] + for p in prompts] + + # Generate + chat_completions = await asyncio.gather( + *(generate_completion(client, model_config, sampling, messages) for messages in batch_messages) + ) + + # Stats + completion_lengths = [len(parse_completion_tokens(c)) for c in chat_completions] + avg_completion_len = sum(completion_lengths) / max(1, len(completion_lengths)) + + completions = parse_completions(chat_completions) + task_types = [item.get("task_type", "") for item in dataset] + verification_infos = [item.get("verification_info", "{}") for item in dataset] + # Duplicate for k samples + task_types = [t for t in task_types for _ in range(rollouts_per_prompt)] + verification_infos = [v for v in verification_infos for _ in range(rollouts_per_prompt)] + + try: + import json + verification_infos = [json.loads(v) if isinstance(v, str) else v for v in verification_infos] + except Exception: + logger.warning("Failed to parse some verification_info entries; using raw values.") + + rewards = compute_rewards(completions, task_types, verification_infos) + + # pass@k for binary rewards + unique = set(rewards) + pass_at_k = None + if unique.issubset({0, 1, 0.0, 1.0}): + k = rollouts_per_prompt + rows: dict[int, list[float]] = {} + for pid, r in zip(problem_ids, rewards): + rows.setdefault(pid, []).append(float(r)) + solved = [any(x == 1.0 for x in rs) for rs in rows.values()] + pass_at_k = sum(solved) / max(1, len(solved)) + + avg_reward = float(sum(map(float, rewards)) / max(1, len(rewards))) + return avg_reward, pass_at_k, float(avg_completion_len) + + +async def score_prompt( + system_prompt: str, + client_cfg: ClientConfig, + model_cfg: OrchestratorModelConfig, + benchmark: str, + subset_size: int, + rollouts_per_prompt: int, + max_tokens: int | None, + min_tokens: int, +) -> PromptScore: + logger = get_logger() + client = setup_client(client_cfg) + + sampling = SamplingConfig( + temperature=1.0, + max_tokens=max_tokens, + min_tokens=min_tokens, + seed=None, + ) + + # Single benchmark specified by caller + avg_reward, pass_at_k, avg_len = await _score_on_benchmark( + client, + benchmark=benchmark, + system_prompt=system_prompt, + model_config=model_cfg, + sampling=sampling, + subset_size=subset_size, + rollouts_per_prompt=rollouts_per_prompt, + ) + + return PromptScore( + prompt=system_prompt, + avg_reward=avg_reward, + pass_at_k=pass_at_k, + avg_completion_len=avg_len, + meta={}, + ) + + +async def score_prompt_dry_run(system_prompt: str, subset_size: int) -> PromptScore: + # Deterministic-ish fake scoring for local CPU testing without dependencies + base = sum(ord(c) for c in system_prompt) % 1000 / 1000.0 + avg_reward = 0.3 + 0.5 * base + pass_at_k = None + avg_completion_len = 20.0 + (len(system_prompt) % 17) + return PromptScore(system_prompt, avg_reward, pass_at_k, avg_completion_len, meta={"dry_run": True}) + + +async def score_prompt_instances( + system_prompt: str, + client_cfg: ClientConfig, + model_cfg: OrchestratorModelConfig, + benchmark: str, + num_instances: int, + rollouts_per_prompt: int, + max_tokens: int | None, + min_tokens: int, + offset: int = 0, +) -> list[float]: + """Return per-instance average reward for a contiguous slice of the dataset.""" + logger = get_logger() + client = setup_client(client_cfg) + sampling = SamplingConfig( + temperature=1.0, + max_tokens=max_tokens, + min_tokens=min_tokens, + seed=None, + ) + + dataset = get_benchmark_dataset(benchmark) + start = max(0, offset) + end = min(len(dataset), start + num_instances) + dataset = dataset.select(range(start, end)) + + prompts = [item["prompt"] for item in dataset] + prompts = [p for p in prompts for _ in range(rollouts_per_prompt)] + problem_ids = list(range(len(dataset))) + problem_ids = [pid for pid in problem_ids for _ in range(rollouts_per_prompt)] + batch_messages = [[{"role": "system", "content": system_prompt}, {"role": "user", "content": p}] for p in prompts] + + chat_completions = await asyncio.gather( + *(generate_completion(client, model_cfg, sampling, messages) for messages in batch_messages) + ) + completions = parse_completions(chat_completions) + task_types = [item.get("task_type", "") for item in dataset] + verification_infos = [item.get("verification_info", "{}") for item in dataset] + task_types = [t for t in task_types for _ in range(rollouts_per_prompt)] + try: + import json + verification_infos = [json.loads(v) if isinstance(v, str) else v for v in verification_infos] + except Exception: + pass + verification_infos = [v for v in verification_infos for _ in range(rollouts_per_prompt)] + rewards = compute_rewards(completions, task_types, verification_infos) + + per_problem: dict[int, list[float]] = {} + for pid, r in zip(problem_ids, rewards): + per_problem.setdefault(pid, []).append(float(r)) + return [sum(rs) / len(rs) for _, rs in sorted(per_problem.items())] + + +async def score_prompt_instances_dry_run(system_prompt: str, num_instances: int) -> list[float]: + base = (sum(ord(c) for c in system_prompt) % 1000) / 1000.0 + return [0.3 + 0.5 * ((base + i * 0.013) % 1.0) for i in range(num_instances)] + + diff --git a/src/prime_rl/optimizer/gepa/gepa.py b/src/prime_rl/optimizer/gepa/gepa.py new file mode 100644 index 0000000000..d0395a9c62 --- /dev/null +++ b/src/prime_rl/optimizer/gepa/gepa.py @@ -0,0 +1,255 @@ +from __future__ import annotations + +import asyncio +import json +import random +from pathlib import Path +from typing import Any + +from prime_rl.optimizer.gepa.config import GEPAConfig +from prime_rl.optimizer.gepa.evaluate import ( + PromptScore, + score_prompt, + score_prompt_dry_run, + score_prompt_instances, + score_prompt_instances_dry_run, +) +from prime_rl.optimizer.gepa.operators import crossover, mutate +from prime_rl.optimizer.gepa.reflection import reflect +from prime_rl.optimizer.gepa.selection import select_indices +from prime_rl.orchestrator.config import ClientConfig, ModelConfig as OrchestratorModelConfig +from prime_rl.utils.monitor import setup_monitor +from prime_rl.utils.pydantic_config import parse_argv +from prime_rl.utils.logger import format_message, format_time, set_logger, setup_handlers, get_logger +from prime_rl.utils.config import LogConfig +from prime_rl.utils.utils import clean_exit + + +async def _evaluate_population( + population: list[str], + cfg: GEPAConfig, +) -> list[PromptScore]: + if cfg.dry_run: + tasks = [score_prompt_dry_run(p, cfg.evaluate.subset_size) for p in population] + return await asyncio.gather(*tasks) + else: + client_cfg = cfg.client + model_cfg = OrchestratorModelConfig(name=cfg.model.name) + tasks = [ + score_prompt( + system_prompt=p, + client_cfg=client_cfg, + model_cfg=model_cfg, + benchmark=cfg.evaluate.benchmark, + subset_size=cfg.evaluate.pareto_size, + rollouts_per_prompt=cfg.evaluate.rollouts_per_prompt, + max_tokens=cfg.evaluate.max_tokens, + min_tokens=cfg.evaluate.min_tokens, + ) + for p in population + ] + return await asyncio.gather(*tasks) + + +def _default_seed_population(base_prompt: str, n: int) -> list[str]: + variants = [ + base_prompt, + base_prompt + "\nBe concise and avoid unnecessary verbosity.", + base_prompt + "\nUse tags to present only the final answer.", + base_prompt + "\nThink step-by-step inside ... before answering.", + ] + while len(variants) < n: + variants.append(base_prompt) + return variants[:n] + + +async def gepa(cfg: GEPAConfig) -> None: + # Setup logger + _setup_logger(cfg.log) + logger = get_logger() + rng = random.Random(cfg.seed) + monitor = setup_monitor(cfg.monitor, outputs_dir=cfg.outputs_dir, run_config=cfg) + + run_dir = cfg.outputs_dir / (cfg.run_name or "gepa") + run_dir.mkdir(parents=True, exist_ok=True) + + # Seed population from reverse-text system prompt as reasonable default baseline + base_prompt = ( + "Follow the task precisely. Use ... for internal reasoning. " + "Output only the final answer inside ...." + ) + population = _default_seed_population(base_prompt, cfg.population_size) + + hall_of_fame: list[tuple[PromptScore, int]] = [] + + rollouts_used = 0 + for gen in range(cfg.generations): + logger.info(f"[GEPA] Generation {gen}: evaluating {len(population)} prompts") + results = await _evaluate_population(population, cfg) + + # Persist population snapshot + snap_path = run_dir / f"gen_{gen}.json" + with open(snap_path, "w") as f: + json.dump([r.__dict__ for r in results], f) + + # Log summary + avg_scores = [r.avg_reward for r in results] + best_idx = max(range(len(results)), key=lambda i: results[i].avg_reward) + best = results[best_idx] + logger.success( + f"[GEPA] Gen {gen}: best={best.avg_reward:.3f} pass@k={best.pass_at_k} len={best.avg_completion_len:.1f}" + ) + monitor.log({ + "gepa/gen": gen, + "gepa/best": best.avg_reward, + "gepa/avg": sum(avg_scores) / max(1, len(avg_scores)), + "step": gen, + }) + + # Update Hall of Fame + hall_of_fame.append((best, gen)) + hall_of_fame = sorted(hall_of_fame, key=lambda x: x[0].avg_reward, reverse=True)[: cfg.selection.keep_elite] + + # Early exit on final gen + if gen == cfg.generations - 1: + break + + # Compute Pareto scores matrix S for D_pareto + if cfg.dry_run: + S = [ + await score_prompt_instances_dry_run(p, cfg.evaluate.pareto_size) + for p in population + ] + else: + S = [] + for p in population: + S.append( + await score_prompt_instances( + p, + cfg.client, + OrchestratorModelConfig(name=cfg.model.name), + cfg.evaluate.benchmark, + cfg.evaluate.pareto_size, + cfg.evaluate.rollouts_per_prompt, + cfg.evaluate.max_tokens, + cfg.evaluate.min_tokens, + offset=0, + ) + ) + + # Pareto-front selection: choose one survivor index + def pareto_select_index(scores: list[list[float]]) -> int: + import math + num_c = len(scores) + num_i = len(scores[0]) if scores else 0 + if num_c == 0 or num_i == 0: + return 0 + # For each instance, find max value and candidates achieving it + max_per_i = [max(scores[c][i] for c in range(num_c)) for i in range(num_i)] + top_sets = [ + {c for c in range(num_c) if math.isclose(scores[c][i], max_per_i[i], rel_tol=1e-9)} + for i in range(num_i) + ] + union = set().union(*top_sets) + # Non-dominated filtering within union + def dominates(a: int, b: int) -> bool: + ge_all = all(scores[a][i] >= scores[b][i] for i in range(num_i)) + gt_any = any(scores[a][i] > scores[b][i] for i in range(num_i)) + return ge_all and gt_any + non_dominated = set(union) + for a in list(union): + for b in list(union): + if a != b and dominates(b, a) and a in non_dominated: + non_dominated.remove(a) + # Coverage weights f: how many instances where candidate is top + cover = {c: 0 for c in non_dominated} + for i, tops in enumerate(top_sets): + for c in non_dominated: + if c in tops: + cover[c] += 1 + # Weighted random choice over non_dominated by coverage + total = sum(cover.values()) or 1 + r = rng.random() * total + acc = 0 + for c, w in cover.items(): + acc += w + if r <= acc: + return c + return next(iter(non_dominated)) + + parent_idx = pareto_select_index(S) + survivors = [parent_idx] + + # Failures: placeholder summaries (TODO: use μ_f feedback) + failures: list[str] = [] + if best.avg_reward < 1.0: + failures.append("answers deviated from correct format or content") + + # Produce next generation with minibatch acceptance gate + next_population: list[str] = [] + # Elites + for score, _gen in hall_of_fame: + next_population.append(score.prompt) + + # Breed survivors + while len(next_population) < cfg.population_size and rollouts_used < cfg.budget_rollouts: + if rng.random() < cfg.operators.crossover_rate and len(survivors) >= 2: + a_idx, b_idx = rng.sample(survivors, 2) + child = crossover(population[a_idx], population[b_idx], cfg.operators, rng) + else: + a_idx = rng.choice(survivors) + child = population[a_idx] + # Reflect then mutate + child = reflect(child, failures, cfg.operators, rng) + if rng.random() < cfg.operators.mutation_rate: + child = mutate(child, cfg.operators, rng) + # Acceptance: compare on a minibatch from feedback pool + if cfg.dry_run: + parent_scores = await score_prompt_instances_dry_run(population[a_idx], cfg.minibatch_size) + child_scores = await score_prompt_instances_dry_run(child, cfg.minibatch_size) + else: + parent_scores = await score_prompt_instances( + population[a_idx], cfg.client, OrchestratorModelConfig(name=cfg.model.name), + cfg.evaluate.benchmark, cfg.minibatch_size, cfg.evaluate.rollouts_per_prompt, + cfg.evaluate.max_tokens, cfg.evaluate.min_tokens, offset=cfg.evaluate.pareto_size, + ) + child_scores = await score_prompt_instances( + child, cfg.client, OrchestratorModelConfig(name=cfg.model.name), + cfg.evaluate.benchmark, cfg.minibatch_size, cfg.evaluate.rollouts_per_prompt, + cfg.evaluate.max_tokens, cfg.evaluate.min_tokens, offset=cfg.evaluate.pareto_size, + ) + + rollouts_used += cfg.minibatch_size + if sum(child_scores) / len(child_scores) > sum(parent_scores) / len(parent_scores): + next_population.append(child) + + population = next_population[: cfg.population_size] + + if rollouts_used >= cfg.budget_rollouts: + logger.info("[GEPA] Budget exhausted; stopping evolution") + break + + # Export best prompt + best_overall = hall_of_fame[0][0] if hall_of_fame else results[0] + out_path = run_dir / "best_prompt.txt" + with open(out_path, "w") as f: + f.write(best_overall.prompt) + logger.success(f"[GEPA] Exported best prompt to {out_path}") + + +def main(): + asyncio.run(gepa(parse_argv(GEPAConfig))) + + +if __name__ == "__main__": + main() + + +def _setup_logger(log_config: LogConfig): + if get_logger(): + return + fmt = format_time(log_config) + format_message() + logger = setup_handlers(__import__("loguru").logger, fmt, log_config, rank=0) + set_logger(logger) + + diff --git a/src/prime_rl/optimizer/gepa/operators.py b/src/prime_rl/optimizer/gepa/operators.py new file mode 100644 index 0000000000..0054ac20cf --- /dev/null +++ b/src/prime_rl/optimizer/gepa/operators.py @@ -0,0 +1,57 @@ +from __future__ import annotations + +import random +import re +from typing import Callable + +from prime_rl.optimizer.gepa.config import OperatorsConfig + + +def _truncate(text: str, max_chars: int) -> str: + return text if len(text) <= max_chars else text[:max_chars] + + +def op_tighten_instruction(prompt: str) -> str: + return re.sub(r"(?i)please |kindly ", "", prompt) + + +def op_enforce_format_tags(prompt: str) -> str: + if "" not in prompt: + prompt += "\nAlways put the final answer inside ...." + return prompt + + +def op_add_reasoning_hint(prompt: str) -> str: + if "" not in prompt and "" not in prompt: + prompt += "\nThink step-by-step inside ... before writing the final answer." + return prompt + + +def op_remove_vagueness(prompt: str) -> str: + prompt = re.sub(r"(?i)try to|attempt to|maybe|possibly|could you ", "", prompt) + return prompt + + +def mutate(prompt: str, cfg: OperatorsConfig, rng: random.Random) -> str: + ops: list[Callable[[str], str]] = [ + op_tighten_instruction, + op_enforce_format_tags, + op_add_reasoning_hint, + op_remove_vagueness, + ] + num_ops = 1 + (rng.random() < 0.5) + for _ in range(num_ops): + op = rng.choice(ops) + prompt = op(prompt) + return _truncate(prompt, cfg.max_prompt_chars) + + +def crossover(a: str, b: str, cfg: OperatorsConfig, rng: random.Random) -> str: + if len(a) < 16 or len(b) < 16: + return a if rng.random() < 0.5 else b + ia = rng.randrange(len(a) // 4, 3 * len(a) // 4) + ib = rng.randrange(len(b) // 4, 3 * len(b) // 4) + child = a[:ia] + "\n" + b[ib:] + return _truncate(child, cfg.max_prompt_chars) + + diff --git a/src/prime_rl/optimizer/gepa/reflection.py b/src/prime_rl/optimizer/gepa/reflection.py new file mode 100644 index 0000000000..1a3696435d --- /dev/null +++ b/src/prime_rl/optimizer/gepa/reflection.py @@ -0,0 +1,32 @@ +from __future__ import annotations + +import random +from typing import Sequence + +from prime_rl.optimizer.gepa.config import OperatorsConfig + + +REFLECT_TEMPLATE = ( + "You are optimizing a system instruction for a model on a benchmark.\n" + "Given a list of failure examples (short summaries), propose concise edits to the instruction\n" + "that will improve accuracy without adding verbosity or vagueness.\n" + "Return only the edited instruction text. Keep it under the specified character limit." +) + + +def reflect(prompt: str, failures: Sequence[str], cfg: OperatorsConfig, rng: random.Random) -> str: + # Placeholder: lightweight heuristic reflection combining failures into a short clause + if not failures: + return prompt + clause = "; ".join(failures[:3]) + edited = ( + prompt + + "\nConstraints: Avoid previous mistakes such as: " + + clause + + ". Be explicit, deterministic, and adhere to required answer format." + ) + if len(edited) > cfg.max_prompt_chars: + edited = edited[: cfg.max_prompt_chars] + return edited + + diff --git a/src/prime_rl/optimizer/gepa/selection.py b/src/prime_rl/optimizer/gepa/selection.py new file mode 100644 index 0000000000..ee1474945f --- /dev/null +++ b/src/prime_rl/optimizer/gepa/selection.py @@ -0,0 +1,39 @@ +from __future__ import annotations + +import random +from typing import Iterable + +from prime_rl.optimizer.gepa.config import SelectionConfig + + +def top_k(scores: list[tuple[int, float]], k: int) -> list[int]: + return [i for i, _ in sorted(scores, key=lambda x: x[1], reverse=True)[:k]] + + +def tournament(scores: list[tuple[int, float]], k: int, t_size: int, rng: random.Random) -> list[int]: + winners: list[int] = [] + indices = [i for i, _ in scores] + for _ in range(k): + pool = rng.sample(indices, min(t_size, len(indices))) + pool_scores = [(i, dict(scores)[i]) for i in pool] + winners.append(max(pool_scores, key=lambda x: x[1])[0]) + return winners + + +def select_indices(cfg: SelectionConfig, fitness: Iterable[float], rng: random.Random) -> list[int]: + scores = list(enumerate(fitness)) + if cfg.strategy == "top-k": + idxs = top_k(scores, cfg.k) + else: + idxs = tournament(scores, cfg.k, cfg.tournament_size, rng) + # Add elites (ensure uniqueness, preserve order) + elites = top_k(scores, cfg.keep_elite) + seen = set() + ordered = [] + for i in elites + idxs: + if i not in seen: + seen.add(i) + ordered.append(i) + return ordered + + diff --git a/src/prime_rl/rl.py b/src/prime_rl/rl.py index 889d24854c..ab8bea782a 100644 --- a/src/prime_rl/rl.py +++ b/src/prime_rl/rl.py @@ -176,6 +176,9 @@ class RLConfig(BaseSettings): ), ] = False + # Optional: run GEPA instead of RL (reuses shared top-level runner UX) + gepa: Annotated[bool, Field(description="If true, run GEPA optimizer instead of RL trainer.")] = False + @model_validator(mode="after") def validate_device(self): available_gpus = torch.cuda.device_count() @@ -406,6 +409,16 @@ def rl(config: RLConfig): logger.info("Starting RL run") logger.debug(f"RL start command: {' '.join(start_command)}") + # Redirect to GEPA optimizer if requested + if getattr(config, "gepa", False): + from prime_rl.optimizer.gepa.gepa import gepa as run_gepa + from prime_rl.optimizer.gepa.config import GEPAConfig + from prime_rl.utils.pydantic_config import parse_argv as parse_gepa_argv + logger.info("GEPA flag set; launching GEPA optimizer") + import asyncio + asyncio.run(run_gepa(parse_gepa_argv(GEPAConfig))) + return + # Prepare paths to communicate with the trainer log_dir = get_log_dir(config.outputs_dir) ckpt_dir = get_ckpt_dir(config.outputs_dir) From 500d7060db325c2bcde5ecfd9e6f23e271c43aab Mon Sep 17 00:00:00 2001 From: MagellaX Date: Thu, 14 Aug 2025 18:11:55 +0530 Subject: [PATCH 2/4] =?UTF-8?q?feat(gepa):=20add=20LLM-based=20reflection?= =?UTF-8?q?=20(=CE=BC=5Ff=20scaffold)=20and=20extend=20env=20system=5Fprom?= =?UTF-8?q?pt=20overrides?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../vf_acereason_math/vf_acereason_math.py | 7 +++- .../vf_deepscaler_math/vf_deepscaler_math.py | 7 +++- .../vf_intellect_math/vf_intellect_math.py | 7 +++- .../vf_pydantic_adherence.py | 8 +--- .../vf_skywork_math/vf_skywork_math.py | 7 +++- src/prime_rl/optimizer/gepa/gepa.py | 18 +++++++-- src/prime_rl/optimizer/gepa/reflection.py | 37 +++++++++++++++++++ 7 files changed, 78 insertions(+), 13 deletions(-) diff --git a/environments/vf_acereason_math/vf_acereason_math.py b/environments/vf_acereason_math/vf_acereason_math.py index f4462fc288..1b6875b654 100644 --- a/environments/vf_acereason_math/vf_acereason_math.py +++ b/environments/vf_acereason_math/vf_acereason_math.py @@ -8,6 +8,7 @@ def load_environment( solve_rate_field: str | None = None, min_solve_rate: float | None = None, max_solve_rate: float | None = None, + system_prompt: str | None = None, **kwargs, ) -> vf.Environment: train_dataset = load_dataset("nvidia/AceReason-Math", split="train").map( @@ -30,5 +31,9 @@ def correct_answer_reward_func(completion, info, **kwargs) -> float: weights=[1.0], ) - vf_env = vf.SingleTurnEnv(dataset=train_dataset, rubric=rubric) + vf_env = ( + vf.SingleTurnEnv(dataset=train_dataset, rubric=rubric, system_prompt=system_prompt) + if system_prompt is not None + else vf.SingleTurnEnv(dataset=train_dataset, rubric=rubric) + ) return vf_env diff --git a/environments/vf_deepscaler_math/vf_deepscaler_math.py b/environments/vf_deepscaler_math/vf_deepscaler_math.py index 7aa77dc75d..81e78639d4 100644 --- a/environments/vf_deepscaler_math/vf_deepscaler_math.py +++ b/environments/vf_deepscaler_math/vf_deepscaler_math.py @@ -8,6 +8,7 @@ def load_environment( solve_rate_field: str | None = None, min_solve_rate: float | None = None, max_solve_rate: float | None = None, + system_prompt: str | None = None, **kwargs, ) -> vf.Environment: train_dataset = load_dataset("agentica-org/DeepScaleR-Preview-Dataset", split="train").map( @@ -30,5 +31,9 @@ def correct_answer_reward_func(completion, info, **kwargs) -> float: weights=[1.0], ) - vf_env = vf.SingleTurnEnv(dataset=train_dataset, rubric=rubric) + vf_env = ( + vf.SingleTurnEnv(dataset=train_dataset, rubric=rubric, system_prompt=system_prompt) + if system_prompt is not None + else vf.SingleTurnEnv(dataset=train_dataset, rubric=rubric) + ) return vf_env diff --git a/environments/vf_intellect_math/vf_intellect_math.py b/environments/vf_intellect_math/vf_intellect_math.py index c5a1c36f1d..fa5a325477 100644 --- a/environments/vf_intellect_math/vf_intellect_math.py +++ b/environments/vf_intellect_math/vf_intellect_math.py @@ -6,6 +6,7 @@ def load_environment( solve_rate_field: str | None = None, min_solve_rate: float | None = None, max_solve_rate: float | None = None, + system_prompt: str | None = None, **kwargs, ) -> vf.Environment: import json @@ -37,5 +38,9 @@ def correct_answer_reward_func(completion, info, **kwargs) -> float: weights=[1.0], ) - vf_env = vf.SingleTurnEnv(dataset=train_dataset, rubric=rubric) + vf_env = ( + vf.SingleTurnEnv(dataset=train_dataset, rubric=rubric, system_prompt=system_prompt) + if system_prompt is not None + else vf.SingleTurnEnv(dataset=train_dataset, rubric=rubric) + ) return vf_env diff --git a/environments/vf_pydantic_adherence/vf_pydantic_adherence.py b/environments/vf_pydantic_adherence/vf_pydantic_adherence.py index 972e02dc12..4ff01b3c77 100644 --- a/environments/vf_pydantic_adherence/vf_pydantic_adherence.py +++ b/environments/vf_pydantic_adherence/vf_pydantic_adherence.py @@ -9,7 +9,7 @@ from verifiers import Messages, Parser -def load_environment() -> vf.Environment: +def load_environment(system_prompt: str | None = None) -> vf.Environment: """ Loads a custom environment. """ @@ -169,10 +169,6 @@ def pydantic_adherence_reward_func(completion, answer, **kwargs): weights=[1.0], ) - vf_env = vf.SingleTurnEnv( - dataset=dataset, - parser=parser, - rubric=rubric, - ) + vf_env = vf.SingleTurnEnv(dataset=dataset, parser=parser, rubric=rubric, system_prompt=system_prompt) if system_prompt is not None else vf.SingleTurnEnv(dataset=dataset, parser=parser, rubric=rubric) return vf_env diff --git a/environments/vf_skywork_math/vf_skywork_math.py b/environments/vf_skywork_math/vf_skywork_math.py index 4c867fc2fb..fd32ed6a7f 100644 --- a/environments/vf_skywork_math/vf_skywork_math.py +++ b/environments/vf_skywork_math/vf_skywork_math.py @@ -10,6 +10,7 @@ def load_environment( solve_rate_field: str | None = None, min_solve_rate: float | None = None, max_solve_rate: float | None = None, + system_prompt: str | None = None, **kwargs, ) -> vf.Environment: train_dataset = load_dataset("PrimeIntellect/Skywork-OR1-RL-Data-v1-math-prime-rl-format", split="train").map( @@ -37,5 +38,9 @@ def correct_answer_reward_func(completion, info, **kwargs) -> float: weights=[1.0], ) - vf_env = vf.SingleTurnEnv(dataset=train_dataset, rubric=rubric) + vf_env = ( + vf.SingleTurnEnv(dataset=train_dataset, rubric=rubric, system_prompt=system_prompt) + if system_prompt is not None + else vf.SingleTurnEnv(dataset=train_dataset, rubric=rubric) + ) return vf_env diff --git a/src/prime_rl/optimizer/gepa/gepa.py b/src/prime_rl/optimizer/gepa/gepa.py index d0395a9c62..79c3bb0016 100644 --- a/src/prime_rl/optimizer/gepa/gepa.py +++ b/src/prime_rl/optimizer/gepa/gepa.py @@ -15,7 +15,7 @@ score_prompt_instances_dry_run, ) from prime_rl.optimizer.gepa.operators import crossover, mutate -from prime_rl.optimizer.gepa.reflection import reflect +from prime_rl.optimizer.gepa.reflection import reflect, reflect_llm from prime_rl.optimizer.gepa.selection import select_indices from prime_rl.orchestrator.config import ClientConfig, ModelConfig as OrchestratorModelConfig from prime_rl.utils.monitor import setup_monitor @@ -199,8 +199,20 @@ def dominates(a: int, b: int) -> bool: else: a_idx = rng.choice(survivors) child = population[a_idx] - # Reflect then mutate - child = reflect(child, failures, cfg.operators, rng) + # Reflect then mutate (use LLM reflection if not dry_run) + if cfg.dry_run: + child = reflect(child, failures, cfg.operators, rng) + else: + try: + child = await reflect_llm( + cfg.client, + OrchestratorModelConfig(name=cfg.model.name), + child, + list(failures), + cfg.operators.max_prompt_chars, + ) + except Exception: + child = reflect(child, failures, cfg.operators, rng) if rng.random() < cfg.operators.mutation_rate: child = mutate(child, cfg.operators, rng) # Acceptance: compare on a minibatch from feedback pool diff --git a/src/prime_rl/optimizer/gepa/reflection.py b/src/prime_rl/optimizer/gepa/reflection.py index 1a3696435d..c08354d871 100644 --- a/src/prime_rl/optimizer/gepa/reflection.py +++ b/src/prime_rl/optimizer/gepa/reflection.py @@ -4,6 +4,8 @@ from typing import Sequence from prime_rl.optimizer.gepa.config import OperatorsConfig +from openai import AsyncOpenAI +from prime_rl.orchestrator.config import ModelConfig as OrchestratorModelConfig, ClientConfig REFLECT_TEMPLATE = ( @@ -30,3 +32,38 @@ def reflect(prompt: str, failures: Sequence[str], cfg: OperatorsConfig, rng: ran return edited +async def reflect_llm( + client_cfg: ClientConfig, + model_cfg: OrchestratorModelConfig, + current_prompt: str, + feedbacks: list[str], + max_chars: int, +) -> str: + """Use the model to propose an edited instruction given feedback examples.""" + client = AsyncOpenAI(base_url=f"http://{client_cfg.host}:{client_cfg.port}/v1", api_key=client_cfg.api_key) + + system = ( + "You are an expert instruction engineer.\n" + "Given the current system instruction and a few failure summaries, propose an improved instruction.\n" + "Constraints: keep it concise, deterministic, preserve required tags/formatting, and stay under the character limit." + ) + fb_text = "\n- ".join([f for f in feedbacks[:5]]) if feedbacks else "(no feedback provided)" + user = ( + f"Current instruction:\n" + f"-----\n{current_prompt}\n-----\n" + f"Failures:\n- {fb_text}\n" + f"Return only the full edited instruction (no commentary), max {max_chars} characters." + ) + resp = await client.chat.completions.create( + model=model_cfg.name, + messages=[{"role": "system", "content": system}, {"role": "user", "content": user}], + temperature=0.2, + max_tokens=512, + ) + text = resp.choices[0].message.content or current_prompt + text = text.strip() + if len(text) > max_chars: + text = text[:max_chars] + return text + + From 70d265a7a433b013f12af26fb584b69be9e9bd9e Mon Sep 17 00:00:00 2001 From: MagellaX Date: Thu, 14 Aug 2025 21:00:34 +0530 Subject: [PATCH 3/4] chore(ruff): fix import ordering and unused imports; quiet unused locals --- src/prime_rl/optimizer/gepa/config.py | 2 +- src/prime_rl/optimizer/gepa/evaluate.py | 15 ++++++++++----- src/prime_rl/optimizer/gepa/gepa.py | 5 ++--- src/prime_rl/optimizer/gepa/reflection.py | 5 +++-- src/prime_rl/rl.py | 2 +- 5 files changed, 17 insertions(+), 12 deletions(-) diff --git a/src/prime_rl/optimizer/gepa/config.py b/src/prime_rl/optimizer/gepa/config.py index ce335e91aa..8e71fc1407 100644 --- a/src/prime_rl/optimizer/gepa/config.py +++ b/src/prime_rl/optimizer/gepa/config.py @@ -3,9 +3,9 @@ from pydantic import Field -from prime_rl.utils.config import LogConfig, MultiMonitorConfig, ModelConfig from prime_rl.inference.config import InferenceConfig from prime_rl.orchestrator.config import ClientConfig +from prime_rl.utils.config import LogConfig, MultiMonitorConfig, ModelConfig from prime_rl.utils.pydantic_config import BaseConfig, BaseSettings diff --git a/src/prime_rl/optimizer/gepa/evaluate.py b/src/prime_rl/optimizer/gepa/evaluate.py index 559458209b..7f01460326 100644 --- a/src/prime_rl/optimizer/gepa/evaluate.py +++ b/src/prime_rl/optimizer/gepa/evaluate.py @@ -8,8 +8,16 @@ from prime_rl.eval.registry import get_benchmark_dataset from prime_rl.orchestrator.client import generate_completion, setup_client -from prime_rl.orchestrator.config import ClientConfig, ModelConfig as OrchestratorModelConfig, SamplingConfig -from prime_rl.orchestrator.utils import compute_rewards, parse_completion_tokens, parse_completions +from prime_rl.orchestrator.config import ( + ClientConfig, + ModelConfig as OrchestratorModelConfig, + SamplingConfig, +) +from prime_rl.orchestrator.utils import ( + compute_rewards, + parse_completion_tokens, + parse_completions, +) from prime_rl.utils.logger import get_logger @@ -71,7 +79,6 @@ async def _score_on_benchmark( unique = set(rewards) pass_at_k = None if unique.issubset({0, 1, 0.0, 1.0}): - k = rollouts_per_prompt rows: dict[int, list[float]] = {} for pid, r in zip(problem_ids, rewards): rows.setdefault(pid, []).append(float(r)) @@ -92,7 +99,6 @@ async def score_prompt( max_tokens: int | None, min_tokens: int, ) -> PromptScore: - logger = get_logger() client = setup_client(client_cfg) sampling = SamplingConfig( @@ -143,7 +149,6 @@ async def score_prompt_instances( offset: int = 0, ) -> list[float]: """Return per-instance average reward for a contiguous slice of the dataset.""" - logger = get_logger() client = setup_client(client_cfg) sampling = SamplingConfig( temperature=1.0, diff --git a/src/prime_rl/optimizer/gepa/gepa.py b/src/prime_rl/optimizer/gepa/gepa.py index 79c3bb0016..8cd59ee4f2 100644 --- a/src/prime_rl/optimizer/gepa/gepa.py +++ b/src/prime_rl/optimizer/gepa/gepa.py @@ -16,11 +16,10 @@ ) from prime_rl.optimizer.gepa.operators import crossover, mutate from prime_rl.optimizer.gepa.reflection import reflect, reflect_llm -from prime_rl.optimizer.gepa.selection import select_indices -from prime_rl.orchestrator.config import ClientConfig, ModelConfig as OrchestratorModelConfig +from prime_rl.orchestrator.config import ModelConfig as OrchestratorModelConfig from prime_rl.utils.monitor import setup_monitor from prime_rl.utils.pydantic_config import parse_argv -from prime_rl.utils.logger import format_message, format_time, set_logger, setup_handlers, get_logger +from prime_rl.utils.logger import format_message, format_time, get_logger, set_logger, setup_handlers from prime_rl.utils.config import LogConfig from prime_rl.utils.utils import clean_exit diff --git a/src/prime_rl/optimizer/gepa/reflection.py b/src/prime_rl/optimizer/gepa/reflection.py index c08354d871..8bde3edde2 100644 --- a/src/prime_rl/optimizer/gepa/reflection.py +++ b/src/prime_rl/optimizer/gepa/reflection.py @@ -3,9 +3,10 @@ import random from typing import Sequence -from prime_rl.optimizer.gepa.config import OperatorsConfig from openai import AsyncOpenAI -from prime_rl.orchestrator.config import ModelConfig as OrchestratorModelConfig, ClientConfig + +from prime_rl.optimizer.gepa.config import OperatorsConfig +from prime_rl.orchestrator.config import ClientConfig, ModelConfig as OrchestratorModelConfig REFLECT_TEMPLATE = ( diff --git a/src/prime_rl/rl.py b/src/prime_rl/rl.py index ab8bea782a..6cb965786c 100644 --- a/src/prime_rl/rl.py +++ b/src/prime_rl/rl.py @@ -24,7 +24,7 @@ from prime_rl.trainer.rl.config import FakeDataLoaderConfig from prime_rl.trainer.rl.config import RLTrainerConfig as TrainerConfig from prime_rl.utils.config import WandbMonitorConfig -from prime_rl.utils.logger import format_message, format_time, get_logger, set_logger, setup_handlers +from prime_rl.utils.logger import get_logger from prime_rl.utils.pydantic_config import BaseSettings, get_temp_toml_file, parse_argv from prime_rl.utils.utils import ( get_ckpt_dir, From 050b57639487b2acfc1d1a8ec20e08e5a281af8d Mon Sep 17 00:00:00 2001 From: MagellaX Date: Fri, 15 Aug 2025 11:50:13 +0530 Subject: [PATCH 4/4] chore: fix rl logger imports and ruff orderings for CI --- src/prime_rl/optimizer/gepa/config.py | 2 +- src/prime_rl/optimizer/gepa/evaluate.py | 4 +++- src/prime_rl/optimizer/gepa/gepa.py | 10 +++++----- src/prime_rl/optimizer/gepa/reflection.py | 4 ++-- src/prime_rl/rl.py | 10 ++++++++-- 5 files changed, 19 insertions(+), 11 deletions(-) diff --git a/src/prime_rl/optimizer/gepa/config.py b/src/prime_rl/optimizer/gepa/config.py index 8e71fc1407..28f3cdf446 100644 --- a/src/prime_rl/optimizer/gepa/config.py +++ b/src/prime_rl/optimizer/gepa/config.py @@ -5,7 +5,7 @@ from prime_rl.inference.config import InferenceConfig from prime_rl.orchestrator.config import ClientConfig -from prime_rl.utils.config import LogConfig, MultiMonitorConfig, ModelConfig +from prime_rl.utils.config import LogConfig, ModelConfig, MultiMonitorConfig from prime_rl.utils.pydantic_config import BaseConfig, BaseSettings diff --git a/src/prime_rl/optimizer/gepa/evaluate.py b/src/prime_rl/optimizer/gepa/evaluate.py index 7f01460326..3ef6d5c1a4 100644 --- a/src/prime_rl/optimizer/gepa/evaluate.py +++ b/src/prime_rl/optimizer/gepa/evaluate.py @@ -10,9 +10,11 @@ from prime_rl.orchestrator.client import generate_completion, setup_client from prime_rl.orchestrator.config import ( ClientConfig, - ModelConfig as OrchestratorModelConfig, SamplingConfig, ) +from prime_rl.orchestrator.config import ( + ModelConfig as OrchestratorModelConfig, +) from prime_rl.orchestrator.utils import ( compute_rewards, parse_completion_tokens, diff --git a/src/prime_rl/optimizer/gepa/gepa.py b/src/prime_rl/optimizer/gepa/gepa.py index 8cd59ee4f2..6171c57f17 100644 --- a/src/prime_rl/optimizer/gepa/gepa.py +++ b/src/prime_rl/optimizer/gepa/gepa.py @@ -3,9 +3,8 @@ import asyncio import json import random -from pathlib import Path -from typing import Any +# (std imports pruned) from prime_rl.optimizer.gepa.config import GEPAConfig from prime_rl.optimizer.gepa.evaluate import ( PromptScore, @@ -17,11 +16,12 @@ from prime_rl.optimizer.gepa.operators import crossover, mutate from prime_rl.optimizer.gepa.reflection import reflect, reflect_llm from prime_rl.orchestrator.config import ModelConfig as OrchestratorModelConfig +from prime_rl.utils.config import LogConfig +from prime_rl.utils.logger import format_message, format_time, get_logger, set_logger, setup_handlers from prime_rl.utils.monitor import setup_monitor from prime_rl.utils.pydantic_config import parse_argv -from prime_rl.utils.logger import format_message, format_time, get_logger, set_logger, setup_handlers -from prime_rl.utils.config import LogConfig -from prime_rl.utils.utils import clean_exit + +# (no clean_exit usage) async def _evaluate_population( diff --git a/src/prime_rl/optimizer/gepa/reflection.py b/src/prime_rl/optimizer/gepa/reflection.py index 8bde3edde2..871ef3a0df 100644 --- a/src/prime_rl/optimizer/gepa/reflection.py +++ b/src/prime_rl/optimizer/gepa/reflection.py @@ -6,8 +6,8 @@ from openai import AsyncOpenAI from prime_rl.optimizer.gepa.config import OperatorsConfig -from prime_rl.orchestrator.config import ClientConfig, ModelConfig as OrchestratorModelConfig - +from prime_rl.orchestrator.config import ClientConfig +from prime_rl.orchestrator.config import ModelConfig as OrchestratorModelConfig REFLECT_TEMPLATE = ( "You are optimizing a system instruction for a model on a benchmark.\n" diff --git a/src/prime_rl/rl.py b/src/prime_rl/rl.py index 6cb965786c..73bd1025ed 100644 --- a/src/prime_rl/rl.py +++ b/src/prime_rl/rl.py @@ -24,7 +24,13 @@ from prime_rl.trainer.rl.config import FakeDataLoaderConfig from prime_rl.trainer.rl.config import RLTrainerConfig as TrainerConfig from prime_rl.utils.config import WandbMonitorConfig -from prime_rl.utils.logger import get_logger +from prime_rl.utils.logger import ( + format_message, + format_time, + get_logger, + set_logger, + setup_handlers, +) from prime_rl.utils.pydantic_config import BaseSettings, get_temp_toml_file, parse_argv from prime_rl.utils.utils import ( get_ckpt_dir, @@ -411,8 +417,8 @@ def rl(config: RLConfig): # Redirect to GEPA optimizer if requested if getattr(config, "gepa", False): - from prime_rl.optimizer.gepa.gepa import gepa as run_gepa from prime_rl.optimizer.gepa.config import GEPAConfig + from prime_rl.optimizer.gepa.gepa import gepa as run_gepa from prime_rl.utils.pydantic_config import parse_argv as parse_gepa_argv logger.info("GEPA flag set; launching GEPA optimizer") import asyncio