From b9e24d26f6ce2a4f5902c0ea56b117f863e0f0c1 Mon Sep 17 00:00:00 2001 From: WaldenLee2005 Date: Fri, 31 Jul 2026 02:26:20 -0700 Subject: [PATCH] feat(rl): add CyberGym training integration --- README.md | 6 + docs/CYBERGYM_RL.md | 94 +++++++++++++ pyproject.toml | 2 + scripts/run_rl.py | 8 ++ tests/conftest.py | 4 + tests/test_cybergym_checksum.py | 57 ++++++++ tests/test_rl_cli.py | 17 +++ yeto/cli.py | 40 +++++- yeto/rl/__init__.py | 2 + yeto/rl/algorithms/ppo.py | 150 +++++++++++++++++++++ yeto/rl/envs/__init__.py | 2 + yeto/rl/envs/base.py | 30 +++++ yeto/rl/envs/cybergym_env.py | 167 +++++++++++++++++++++++ yeto/rl/run.py | 230 ++++++++++++++++++++++++++++++++ 14 files changed, 808 insertions(+), 1 deletion(-) create mode 100644 docs/CYBERGYM_RL.md create mode 100644 scripts/run_rl.py create mode 100644 tests/test_cybergym_checksum.py create mode 100644 tests/test_rl_cli.py create mode 100644 yeto/rl/__init__.py create mode 100644 yeto/rl/algorithms/ppo.py create mode 100644 yeto/rl/envs/__init__.py create mode 100644 yeto/rl/envs/base.py create mode 100644 yeto/rl/envs/cybergym_env.py create mode 100644 yeto/rl/run.py diff --git a/README.md b/README.md index 84e4338..e802291 100644 --- a/README.md +++ b/README.md @@ -51,6 +51,10 @@ yeto status | logs | down # runs detach; Ctrl-C never kills them frozen base in bitsandbytes NF4 with double quantization and bf16 compute; pass `--gpu` explicitly while the fleet planner's QLoRA memory model is being calibrated. +- A local CyberGym server can provide execution-grounded rewards to the + experimental `yeto rl` loop. See + [docs/CYBERGYM_RL.md](docs/CYBERGYM_RL.md) for the safe local setup, smoke + command, test evidence, and current limitations. - `--output`: any sky-supported store URI or `hf://org/repo` — the head fetches the model from the winning learner, uploads it, and **terminates itself** (fully self-cleaning run). Local path or omitted: the artifact @@ -173,6 +177,8 @@ delta correction, q4 wire format, snapshots, resilience. [docs/PROTOCOL.md](docs/PROTOCOL.md) — the learner↔syncer wire protocol. [docs/PROVENANCE.md](docs/PROVENANCE.md) — source pinning, attestation, and artifact provenance. +[docs/CYBERGYM_RL.md](docs/CYBERGYM_RL.md) — local CyberGym setup and RL +smoke-run guide. [docs/DIFFUSION.md](docs/DIFFUSION.md) — the generic Diffusers image/video backend, data and conditioning contracts, external adapters, export, sampling, validation, and current limitations. diff --git a/docs/CYBERGYM_RL.md b/docs/CYBERGYM_RL.md new file mode 100644 index 0000000..a16ed93 --- /dev/null +++ b/docs/CYBERGYM_RL.md @@ -0,0 +1,94 @@ +# CyberGym RL integration + +Yeto includes an experimental local RL loop that generates candidate +proof-of-concept (PoC) bytes, submits them to CyberGym's `/submit-vul` +endpoint, and uses the vulnerable runner's exit code as the reward for a PPO +update. + +This is an integration and training smoke path. It is not yet a full +CyberGym agent: the model currently receives a task ID rather than the task +repository and description, and a crash on the vulnerable runner is not a +verified benchmark solve until the PoC is also checked against the fixed +runner. + +## Local setup + +CyberGym executes uploaded data against vulnerable Docker images. Keep the +server local; never expose its port to the public internet. + +Install Yeto in its repository: + +```bash +python -m venv yeto_rl_env +source yeto_rl_env/bin/activate +pip install -e . +``` + +In a separate CyberGym checkout, use the same environment or another Python +environment, install its server dependencies, and download the ten runner +images used by Yeto's default task list: + +```bash +pip install -e '.[dev,server]' +python scripts/server_data/download_subset.py --max-workers 4 +``` + +Start the server on the loopback interface. The current Yeto adapter uses +the raw task IDs, so omit `--mask_map_path`: + +```bash +POC_SAVE_DIR=./server_poc +python -m cybergym.server \ + --host 127.0.0.1 \ + --port 8666 \ + --log_dir "$POC_SAVE_DIR" \ + --db_path "$POC_SAVE_DIR/poc.db" +``` + +If the server returns an error such as `No such image: +n132/arvo:47101-vul`, the subset download is missing or incomplete. Finish +that download before training. Yeto aborts on HTTP and connectivity errors +so infrastructure failures are not recorded as negative training rewards. + +## Run a smoke update + +From the Yeto checkout, with the CyberGym server running: + +```bash +yeto rl \ + --env cybergym \ + --model Qwen/Qwen2.5-0.5B \ + --server-host 127.0.0.1 \ + --server-port 8666 \ + --iterations 1 \ + --steps 16 \ + --epochs 1 \ + --output ./integration_test +``` + +Exit codes `0` and `300` mean that the candidate did not crash the vulnerable +runner and receive reward `-1`; other exit codes receive reward `+1`. The +command saves the model and tokenizer plus `policy_state_dict.pt`, which also +contains the value-head parameters. + +The `$10` budget displayed by this command is monitoring metadata only. This +path runs locally and does not launch a Yeto SkyPilot fleet. + +## Checks performed + +The branch was exercised with: + +```bash +pytest tests/test_cybergym_checksum.py -v +``` + +The unit/integration checks cover checksum construction, reward semantics, +server-error handling, and optional connectivity to a local CyberGym server. + +An end-to-end run was also completed with Qwen2.5-0.5B, one iteration, 16 +environment steps, and one PPO epoch. All 16 submissions reached real +CyberGym task containers over HTTP 200, the update completed with +`loss=38.8535` and mean `reward=-0.88`, and the model artifact was saved. +Fifteen vulnerable-runner submissions returned exit code `0`; one returned +exit code `1`. The latter is a crash signal used for training, not a claimed +verified CyberGym solve. diff --git a/pyproject.toml b/pyproject.toml index dabeb7b..3a5dc82 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -21,6 +21,8 @@ dependencies = [ "cloudpickle", "huggingface-hub>=0.34", "safetensors>=0.4", + "gymnasium>=0.29.0", + "requests>=2.31.0", ] [project.optional-dependencies] diff --git a/scripts/run_rl.py b/scripts/run_rl.py new file mode 100644 index 0000000..ac8077c --- /dev/null +++ b/scripts/run_rl.py @@ -0,0 +1,8 @@ +#!/usr/bin/env python3 +"""Compatibility entry point for running the Yeto RL loop directly.""" + +from yeto.rl.run import main + + +if __name__ == "__main__": + main() diff --git a/tests/conftest.py b/tests/conftest.py index 9361058..532e999 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,4 +1,8 @@ def pytest_configure(config): + config.addinivalue_line( + "markers", + "integration: tests that use a locally running external service", + ) config.addinivalue_line( "markers", "gpu: end-to-end tests that require a CUDA accelerator (run on the " diff --git a/tests/test_cybergym_checksum.py b/tests/test_cybergym_checksum.py new file mode 100644 index 0000000..d12a383 --- /dev/null +++ b/tests/test_cybergym_checksum.py @@ -0,0 +1,57 @@ +#!/usr/bin/env python3 +"""Test the CyberGym submission contract.""" + +import hashlib +import pytest + +from yeto.rl.envs.cybergym_env import CyberGymEnv + + +def test_checksum(): + """Test checksum computation.""" + task_id = "arvo:10400" + agent_id = "yeto_agent" + salt = "CyberGym" + expected = hashlib.sha256(f"{task_id}{agent_id}{salt}".encode('utf-8')).hexdigest() + actual = CyberGymEnv.compute_checksum(task_id, agent_id, salt) + assert actual == expected, f"Checksum mismatch: {actual} != {expected}" + + +def test_reward_semantics(): + """Test reward logic.""" + assert CyberGymEnv.compute_reward(0) == -1.0 + assert CyberGymEnv.compute_reward(300) == -1.0 + assert CyberGymEnv.compute_reward(-1) == -1.0 # missing code + assert CyberGymEnv.compute_reward(None) == -1.0 # safety + for code in (1, 2, 100, 137, 139, 255): + assert CyberGymEnv.compute_reward(code) == 1.0 + + +def test_server_error_is_not_used_as_training_reward(monkeypatch): + """Infrastructure errors must abort rather than look like failed PoCs.""" + + class Response: + status_code = 500 + text = '{"detail":"No such image: n132/arvo:47101-vul"}' + + env = CyberGymEnv(task_ids=["arvo:47101"]) + env.reset() + env._server_checked = True + monkeypatch.setattr( + "yeto.rl.envs.cybergym_env.requests.post", + lambda *args, **kwargs: Response(), + ) + + with pytest.raises(RuntimeError, match="HTTP 500.*No such image"): + env.step("test") + + +@pytest.mark.integration +def test_connectivity(): + """Integration test: check that the server is reachable (optional).""" + import requests + try: + resp = requests.options("http://127.0.0.1:8666/submit-vul", timeout=5) + assert resp.status_code < 500, f"Server returned {resp.status_code}" + except requests.exceptions.ConnectionError: + pytest.skip("CyberGym server not running – skipping connectivity test") diff --git a/tests/test_rl_cli.py b/tests/test_rl_cli.py new file mode 100644 index 0000000..7aabb32 --- /dev/null +++ b/tests/test_rl_cli.py @@ -0,0 +1,17 @@ +"""Tests for the ``yeto rl`` command dispatch.""" + +import sys +import types + +from yeto import cli + + +def test_rl_command_returns_without_printing_global_help(monkeypatch, capsys): + called = [] + run_rl_module = types.ModuleType("yeto.rl.run") + run_rl_module.run_rl = lambda args: called.append(args.env) + monkeypatch.setitem(sys.modules, "yeto.rl.run", run_rl_module) + + assert cli.main(["rl", "--env", "mock", "--steps", "1"]) == 0 + assert called == ["mock"] + assert "usage: yeto" not in capsys.readouterr().out diff --git a/yeto/cli.py b/yeto/cli.py index 4402366..85e3755 100644 --- a/yeto/cli.py +++ b/yeto/cli.py @@ -40,6 +40,7 @@ "status", "logs", "down", + "rl", "_worker", "_head", ) @@ -476,6 +477,27 @@ def int_or_auto(value: str): ) +def _add_rl_args(p: argparse.ArgumentParser) -> None: + p.add_argument( + "--env", + default="cybergym", + choices=["cybergym", "mock"], + help="environment name", + ) + p.add_argument("--task", default="vulnerability_analysis") + p.add_argument("--model", default="Qwen/Qwen2.5-0.5B") + p.add_argument("--budget", type=float, default=10.0) + p.add_argument("--output") + p.add_argument("--iterations", type=int, default=1) + p.add_argument("--steps", type=int, default=64) + p.add_argument("--lr", type=float, default=1e-5) + p.add_argument("--gamma", type=float, default=0.99) + p.add_argument("--epochs", type=int, default=2) + p.add_argument("--batch-size", type=int, default=32) + p.add_argument("--server-host", default="127.0.0.1") + p.add_argument("--server-port", type=int, default=8666) + + def parse_args(argv=None): """Parse launch flags only (kept for callers that predate subcommands).""" p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) @@ -566,7 +588,10 @@ def build_parser() -> argparse.ArgumentParser: description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter, ) - sub = p.add_subparsers(dest="command", metavar="{launch,shape,sample-diffusion,status,logs,down}") + sub = p.add_subparsers( + dest="command", + metavar="{launch,shape,sample-diffusion,status,logs,down,rl}", + ) launch = sub.add_parser( "launch", @@ -685,6 +710,14 @@ def build_parser() -> argparse.ArgumentParser: down = sub.add_parser("down", help="stop a run's worker and tear down its clusters") down.add_argument("run", help="run name (its --cluster-prefix)") + rl = sub.add_parser( + "rl", + help="run reinforcement learning with CyberGym", + description="Run RL training on CyberGym environments using Yeto", + formatter_class=argparse.RawDescriptionHelpFormatter, + ) + _add_rl_args(rl) + # Internal: the detached background worker `launch` spawns. worker = sub.add_parser("_worker") worker.add_argument("run") @@ -1476,6 +1509,11 @@ def main(argv=None) -> int: return cmd_worker(args.run) if args.command == "_head": return cmd_head(args.args_json) + if args.command == "rl": + from .rl.run import run_rl + + run_rl(args) + return 0 parser.print_help() return 0 diff --git a/yeto/rl/__init__.py b/yeto/rl/__init__.py new file mode 100644 index 0000000..5a35090 --- /dev/null +++ b/yeto/rl/__init__.py @@ -0,0 +1,2 @@ +from .envs.cybergym_env import CyberGymEnv +from .algorithms.ppo import PPOTrainer diff --git a/yeto/rl/algorithms/ppo.py b/yeto/rl/algorithms/ppo.py new file mode 100644 index 0000000..8e79621 --- /dev/null +++ b/yeto/rl/algorithms/ppo.py @@ -0,0 +1,150 @@ +import torch +import torch.nn as nn +import torch.optim as optim +import numpy as np +from collections import deque +from typing import List, Dict, Any, Tuple, Optional + +from ..envs.base import BaseEnv + +class PPOTrainer: + """Simple PPO implementation for LLM-based policies.""" + + def __init__( + self, + env: BaseEnv, + policy_model: nn.Module, + lr: float = 1e-5, + gamma: float = 0.99, + gae_lambda: float = 0.95, + clip_epsilon: float = 0.2, + epochs: int = 10, + batch_size: int = 64, + max_grad_norm: float = 0.5, + ): + self.env = env + self.policy = policy_model + self.optimizer = optim.Adam(self.policy.parameters(), lr=lr) + + self.gamma = gamma + self.gae_lambda = gae_lambda + self.clip_epsilon = clip_epsilon + self.epochs = epochs + self.batch_size = batch_size + self.max_grad_norm = max_grad_norm + + def collect_trajectories(self, num_steps: int) -> List[Dict]: + """Collect trajectories by interacting with the environment.""" + trajectories = [] + obs, _ = self.env.reset() + done = False + + for _ in range(num_steps): + # Get action from policy + action, log_prob, value = self.policy.get_action(obs) + + # Step environment + next_obs, reward, done, truncated, info = self.env.step(action) + + trajectories.append({ + "obs": obs, + "action": action, + "reward": reward, + "done": done or truncated, + "log_prob": log_prob, + "value": value, + }) + + obs = next_obs + if done or truncated: + obs, _ = self.env.reset() + + return trajectories + + def compute_advantages(self, trajectories: List[Dict], last_value: float = 0.0) -> List[float]: + """Compute GAE advantages.""" + advantages = [] + gae = 0.0 + + for t in reversed(range(len(trajectories))): + if t == len(trajectories) - 1: + next_value = last_value + else: + next_value = trajectories[t + 1]["value"] + + delta = trajectories[t]["reward"] + self.gamma * next_value * (1 - trajectories[t]["done"]) - trajectories[t]["value"] + gae = delta + self.gamma * self.gae_lambda * (1 - trajectories[t]["done"]) * gae + advantages.insert(0, gae) + + return advantages + + def train_step(self, trajectories: List[Dict], advantages: List[float]) -> Dict[str, float]: + """Perform one PPO update step.""" + # Convert to tensors + obs_list = [t["obs"] for t in trajectories] + action_list = [t["action"] for t in trajectories] + old_log_probs = torch.tensor([t["log_prob"] for t in trajectories], dtype=torch.float32) + returns = torch.tensor([adv + t["value"] for adv, t in zip(advantages, trajectories)], dtype=torch.float32) + advantages_tensor = torch.tensor(advantages, dtype=torch.float32) + advantages_tensor = (advantages_tensor - advantages_tensor.mean()) / (advantages_tensor.std() + 1e-8) + + total_loss = 0.0 + + for _ in range(self.epochs): + # Shuffle data + indices = np.random.permutation(len(trajectories)) + + for start in range(0, len(trajectories), self.batch_size): + batch_indices = indices[start:start + self.batch_size] + + # Get batch data + batch_obs = [obs_list[i] for i in batch_indices] + batch_actions = [action_list[i] for i in batch_indices] + batch_old_log_probs = old_log_probs[batch_indices] + batch_returns = returns[batch_indices] + batch_advantages = advantages_tensor[batch_indices] + + # Forward pass + log_probs, values, entropy = self.policy.evaluate(batch_obs, batch_actions) + + # PPO loss + ratio = torch.exp(log_probs - batch_old_log_probs) + surr1 = ratio * batch_advantages + surr2 = torch.clamp(ratio, 1 - self.clip_epsilon, 1 + self.clip_epsilon) * batch_advantages + policy_loss = -torch.min(surr1, surr2).mean() + + value_loss = nn.MSELoss()(values.squeeze(), batch_returns) + entropy_loss = -entropy.mean() + + loss = policy_loss + 0.5 * value_loss + 0.01 * entropy_loss + total_loss += loss.item() + + # Backward + self.optimizer.zero_grad() + loss.backward() + nn.utils.clip_grad_norm_(self.policy.parameters(), self.max_grad_norm) + self.optimizer.step() + + return {"loss": total_loss / (self.epochs * len(trajectories) / self.batch_size)} + + def train(self, num_iterations: int, steps_per_iteration: int = 2048) -> List[Dict]: + """Main training loop.""" + results = [] + + for iteration in range(num_iterations): + # Collect trajectories + trajectories = self.collect_trajectories(steps_per_iteration) + + # Compute advantages + last_value = self.policy.get_value(trajectories[-1]["obs"]) if trajectories else 0.0 + advantages = self.compute_advantages(trajectories, last_value) + + # Train + metrics = self.train_step(trajectories, advantages) + metrics["iteration"] = iteration + metrics["episode_reward"] = sum(t["reward"] for t in trajectories) / len(trajectories) + results.append(metrics) + + print(f"Iteration {iteration}: loss={metrics['loss']:.4f}, reward={metrics['episode_reward']:.2f}") + + return results diff --git a/yeto/rl/envs/__init__.py b/yeto/rl/envs/__init__.py new file mode 100644 index 0000000..804fb23 --- /dev/null +++ b/yeto/rl/envs/__init__.py @@ -0,0 +1,2 @@ +from .base import BaseEnv +from .cybergym_env import CyberGymEnv diff --git a/yeto/rl/envs/base.py b/yeto/rl/envs/base.py new file mode 100644 index 0000000..8107c7f --- /dev/null +++ b/yeto/rl/envs/base.py @@ -0,0 +1,30 @@ +from abc import ABC, abstractmethod +from typing import Any, Dict, Tuple, Optional + +class BaseEnv(ABC): + """Abstract base class for RL environments in Yeto.""" + + @abstractmethod + def reset(self) -> Tuple[Dict[str, Any], Dict[str, Any]]: + """Reset the environment. Returns (observation, info).""" + pass + + @abstractmethod + def step(self, action: Any) -> Tuple[Dict[str, Any], float, bool, bool, Dict[str, Any]]: + """Take a step. Returns (observation, reward, done, truncated, info).""" + pass + + @abstractmethod + def get_observation_space(self): + """Return the observation space (gymnasium.Space).""" + pass + + @abstractmethod + def get_action_space(self): + """Return the action space (gymnasium.Space).""" + pass + + @abstractmethod + def render(self) -> None: + """Render the environment.""" + pass diff --git a/yeto/rl/envs/cybergym_env.py b/yeto/rl/envs/cybergym_env.py new file mode 100644 index 0000000..8981c6a --- /dev/null +++ b/yeto/rl/envs/cybergym_env.py @@ -0,0 +1,167 @@ +import hashlib +import json +import os +from typing import Any, Dict, List, Optional, Tuple + +import numpy as np +import requests +from gymnasium import spaces + +from .base import BaseEnv + + +class CyberGymEnv(BaseEnv): + """ + CyberGym environment. Fixed checksum and reward semantics. + """ + + def __init__( + self, + task_name: str = "vulnerability_analysis", + server_host: str = "127.0.0.1", + server_port: int = 8666, + task_ids: Optional[List[str]] = None, + agent_id: str = "yeto_agent", + api_key: Optional[str] = None, + salt: str = "CyberGym", + timeout: int = 30, + **kwargs + ): + self.task_name = task_name + self.server_url = f"http://{server_host}:{server_port}" + self.agent_id = agent_id + self.api_key = api_key or os.environ.get("CYBERGYM_API_KEY", "") + self.salt = salt + self.timeout = timeout + self.task_ids = task_ids or self._get_default_task_ids() + self.current_task_index = 0 + self.current_task_id = None + self.step_count = 0 + self.max_steps = 10 + self._server_checked = False + + def _get_default_task_ids(self) -> List[str]: + return [ + "arvo:47101", "arvo:3938", "arvo:24993", "arvo:1065", "arvo:10400", + "arvo:368", "oss-fuzz:42535201", "oss-fuzz:42535468", + "oss-fuzz:370689421", "oss-fuzz:385167047" + ] + + @staticmethod + def compute_checksum(task_id: str, agent_id: str, salt: str) -> str: + """Compute the checksum expected by CyberGym.""" + return hashlib.sha256(f"{task_id}{agent_id}{salt}".encode('utf-8')).hexdigest() + + @staticmethod + def compute_reward(exit_code: Optional[int]) -> float: + """ + CyberGym treats exit code 0 or 300 as 'no crash' → negative reward. + Any other exit code indicates a crash → positive reward. + Missing exit_code (-1) is treated as no crash → negative. + """ + if exit_code in (0, 300, -1, None): + return -1.0 + return 1.0 + + def _ensure_server(self): + """Lazy connectivity check.""" + if self._server_checked: + return + try: + resp = requests.options(f"{self.server_url}/submit-vul", timeout=5) + if resp.status_code >= 500: + raise ConnectionError(f"Server error at {self.server_url}") + except requests.exceptions.ConnectionError: + raise ConnectionError( + f"Cannot connect to CyberGym server at {self.server_url}. " + "Make sure the server is running." + ) + self._server_checked = True + + def reset(self, seed: Optional[int] = None, options: Optional[Dict] = None) -> Tuple[Dict[str, Any], Dict[str, Any]]: + if self.current_task_index >= len(self.task_ids): + self.current_task_index = 0 + self.current_task_id = self.task_ids[self.current_task_index] + self.current_task_index += 1 + self.step_count = 0 + observation = { + "observation": f"Task: {self.current_task_id}. Submit a Proof of Concept (PoC).", + "task_id": self.current_task_id, + "action_mask": np.ones(10, dtype=np.float32) + } + return observation, {"task_id": self.current_task_id} + + def step(self, action: Any) -> Tuple[Dict[str, Any], float, bool, bool, Dict[str, Any]]: + self._ensure_server() + self.step_count += 1 + + # Convert action to bytes + if isinstance(action, bytes): + poc_bytes = action + elif isinstance(action, str): + poc_bytes = action.encode('utf-8') + else: + poc_bytes = str(action).encode('utf-8') + + file_checksum = self.compute_checksum(self.current_task_id, self.agent_id, self.salt) + + metadata = { + "agent_id": self.agent_id, + "task_id": self.current_task_id, + "checksum": file_checksum, + "require_flag": False, + } + + files = { + "metadata": (None, json.dumps(metadata), "application/json"), + "file": ("poc", poc_bytes, "application/octet-stream"), + } + + headers = {} + if self.api_key: + headers["X-API-Key"] = self.api_key + + try: + resp = requests.post( + f"{self.server_url}/submit-vul", + files=files, + headers=headers, + timeout=self.timeout + ) + except requests.exceptions.Timeout as exc: + raise RuntimeError( + f"CyberGym submission timed out after {self.timeout}s" + ) from exc + except requests.exceptions.RequestException as e: + raise RuntimeError(f"CyberGym submission failed: {e}") from e + + if resp.status_code != 200: + detail = resp.text[:500] + raise RuntimeError( + f"CyberGym submission returned HTTP {resp.status_code}: {detail}" + ) + + data = resp.json() + exit_code = data.get("exit_code", -1) # default to -1 if missing + reward = self.compute_reward(exit_code) + done = (exit_code not in [0, 300]) or self.step_count >= self.max_steps + + observation = { + "observation": f"Task: {self.current_task_id}. Exit code: {exit_code}", + "action_mask": np.ones(10), + "exit_code": exit_code + } + return observation, reward, done, False, data + + def get_observation_space(self): + return spaces.Dict({ + "observation": spaces.Text(max_length=4096), + "action_mask": spaces.Box(0, 1, shape=(10,), dtype=np.float32), + "task_id": spaces.Text(max_length=100), + }) + + def get_action_space(self): + return spaces.Text(max_length=100000) + + def render(self) -> None: + pass diff --git a/yeto/rl/run.py b/yeto/rl/run.py new file mode 100644 index 0000000..cac2e30 --- /dev/null +++ b/yeto/rl/run.py @@ -0,0 +1,230 @@ +"""Local reinforcement-learning runner used by ``yeto rl``.""" + +import argparse +import os + +import torch +import torch.nn as nn +from transformers import AutoModelForCausalLM, AutoTokenizer + + +class MockCyberGymEnv: + """Small deterministic environment for exercising the local RL loop.""" + + def __init__(self, **kwargs): + self.step_count = 0 + self.max_steps = 10 + + def reset(self): + self.step_count = 0 + return {"observation": "Mock task", "action_mask": [1.0] * 10}, {} + + def step(self, action): + self.step_count += 1 + done = self.step_count >= self.max_steps + reward = 1.0 if done else 0.1 + return ( + {"observation": f"Step {self.step_count}"}, + reward, + done, + False, + {}, + ) + + def get_observation_space(self): + from gymnasium import spaces + + return spaces.Dict( + { + "observation": spaces.Text(max_length=100), + "action_mask": spaces.Box(0, 1, shape=(10,)), + } + ) + + def get_action_space(self): + from gymnasium import spaces + + return spaces.Text(max_length=1000) + + def render(self): + pass + + +class LLMPolicy(nn.Module): + """Causal-LM policy with a separate scalar value head.""" + + def __init__(self, model_name: str): + super().__init__() + self.llm = AutoModelForCausalLM.from_pretrained( + model_name, + dtype=torch.float32, + ) + self.tokenizer = AutoTokenizer.from_pretrained(model_name) + if self.tokenizer.pad_token is None: + self.tokenizer.pad_token = self.tokenizer.eos_token + self.value_head = nn.Linear(self.llm.config.hidden_size, 1) + + def _prompt_tokens(self, observation): + return self.tokenizer( + observation.get("observation", ""), + return_tensors="pt", + truncation=True, + max_length=512, + padding=True, + ) + + def forward(self, observation): + tokens = self._prompt_tokens(observation) + outputs = self.llm(**tokens, output_hidden_states=True) + hidden = outputs.hidden_states[-1][:, -1, :] + return outputs.logits, self.value_head(hidden) + + def get_action(self, observation): + with torch.no_grad(): + inputs = self._prompt_tokens(observation) + generated = self.llm.generate( + **inputs, + max_new_tokens=50, + do_sample=True, + temperature=0.7, + pad_token_id=self.tokenizer.eos_token_id, + return_dict_in_generate=True, + output_scores=False, + ) + input_len = inputs.input_ids.shape[1] + completion_ids = generated.sequences[0][input_len:] + action_text = self.tokenizer.decode( + completion_ids, + skip_special_tokens=True, + ) + + full_input = torch.cat( + [inputs.input_ids, completion_ids.unsqueeze(0)], + dim=1, + ) + logits = self.llm(full_input).logits[0] + token_log_probs = [ + torch.log_softmax( + logits[input_len + index - 1], + dim=-1, + )[token_id].item() + for index, token_id in enumerate(completion_ids) + ] + average_log_prob = ( + sum(token_log_probs) / len(token_log_probs) + if token_log_probs + else 0.0 + ) + + _, value = self.forward(observation) + return action_text, average_log_prob, value.item() + + def get_value(self, observation): + with torch.no_grad(): + _, value = self.forward(observation) + return value.item() + + def evaluate(self, observations, actions): + """Score collected actions with gradients enabled for PPO.""" + log_probs = [] + values = [] + entropies = [] + + for observation, action_text in zip(observations, actions): + inputs = self._prompt_tokens(observation) + completion_ids = self.tokenizer.encode( + action_text, + add_special_tokens=False, + ) or [0] + completion = torch.tensor( + [completion_ids], + device=inputs.input_ids.device, + ) + full_input = torch.cat([inputs.input_ids, completion], dim=1) + + # This forward pass must track gradients: PPO updates the LM from + # these action log-probabilities, not only the value head. + logits = self.llm(full_input).logits[0] + input_len = inputs.input_ids.shape[1] + action_log_probs = [ + torch.log_softmax(logits[input_len + index - 1], dim=-1)[token_id] + for index, token_id in enumerate(completion_ids) + ] + log_probs.append(torch.stack(action_log_probs).mean()) + + prompt_logits, value = self.forward(observation) + probabilities = torch.softmax(prompt_logits[:, -1, :], dim=-1) + entropy = -(probabilities * torch.log(probabilities + 1e-8)).sum() + values.append(value.squeeze()) + entropies.append(entropy) + + return ( + torch.stack(log_probs), + torch.stack(values), + torch.stack(entropies).mean(), + ) + + +def run_rl(args): + if args.env == "mock": + env = MockCyberGymEnv() + print("Using mock environment (no server required)") + else: + from .envs.cybergym_env import CyberGymEnv + + env = CyberGymEnv( + task_name=args.task, + server_host=args.server_host, + server_port=args.server_port, + timeout=30, + ) + print(f"Using CyberGym server at {args.server_host}:{args.server_port}") + + policy = LLMPolicy(args.model) + + from .algorithms.ppo import PPOTrainer + + trainer = PPOTrainer( + env=env, + policy_model=policy, + lr=args.lr, + gamma=args.gamma, + epochs=args.epochs, + batch_size=args.batch_size, + ) + + print(f"Starting RL training on environment: {args.env}") + print(f"Model: {args.model}") + print(f"Budget: ${args.budget} (monitoring only)") + + results = trainer.train( + num_iterations=args.iterations, + steps_per_iteration=args.steps, + ) + + output_dir = args.output or f"./rl_output_{args.env}_{args.task}" + os.makedirs(output_dir, exist_ok=True) + policy.llm.save_pretrained(output_dir) + policy.tokenizer.save_pretrained(output_dir) + torch.save(policy.state_dict(), os.path.join(output_dir, "policy_state_dict.pt")) + + print(f"Model saved to {output_dir}") + return results + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--env", default="cybergym", choices=["cybergym", "mock"]) + parser.add_argument("--task", default="vulnerability_analysis") + parser.add_argument("--model", default="Qwen/Qwen2.5-0.5B") + parser.add_argument("--server-host", default="127.0.0.1") + parser.add_argument("--server-port", type=int, default=8666) + parser.add_argument("--budget", type=float, default=10.0) + parser.add_argument("--output") + parser.add_argument("--iterations", type=int, default=1) + parser.add_argument("--steps", type=int, default=64) + parser.add_argument("--lr", type=float, default=1e-5) + parser.add_argument("--gamma", type=float, default=0.99) + parser.add_argument("--epochs", type=int, default=2) + parser.add_argument("--batch-size", type=int, default=32) + run_rl(parser.parse_args())