{ "cells": [ { "cell_type": "markdown", "metadata": {}, "source": [ "# ContextCorruption-Env — GRPO Training\n", "> **OpenEnv Hackathon | Meta × HuggingFace × PyTorch**\n", "\n", "Fine-tunes **Qwen2-1.5B-Instruct** with GRPO to identify corrupted documents and answer questions correctly.\n", "\n", "**Reward signal (fully deterministic, no LLM judge):**\n", "| Component | Weight |\n", "|---|---|\n", "| Answer correctness (exact match after normalisation) | +0.40 |\n", "| Corruption detection recall | +0.30 |\n", "| False-positive penalty | +0.20 |\n", "| Confidence calibration | ±0.10 |\n", "| Efficiency bonus | +0.05 |\n", "\n", "**Random baseline:** avg reward ≈ 0.13 — beat this to show improvement.\n", "\n", "---\n", "⚠️ Requires **GPU runtime** (A100 recommended). Go to `Runtime → Change runtime type → GPU`." ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 1. Install dependencies" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "%%capture\n", "!pip install openenv-core==0.2.3 unsloth trl transformers datasets wandb faker python-dotenv" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 2. Clone repo and generate facts" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "import os\n", "\n", "REPO_URL = \"https://github.com/sas-dev5/context-corruption-env.git\"\n", "\n", "!git clone {REPO_URL}\n", "%cd context-corruption-env\n", "\n", "# Generate facts.json (pulls NQ + PopQA)\n", "!python -m data.loader" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 3. Authenticate WandB and HuggingFace" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "import wandb\n", "from huggingface_hub import login\n", "\n", "# Paste your keys here or set as Colab secrets\n", "WANDB_API_KEY = os.getenv(\"WANDB_API_KEY\", \"\")\n", "HF_TOKEN = os.getenv(\"HF_TOKEN\", \"\")\n", "HF_HUB_MODEL_ID = \"\" # e.g. \"your-username/qwen-1.5b-context-corruption\" — leave blank to skip\n", "\n", "if WANDB_API_KEY:\n", " wandb.login(key=WANDB_API_KEY)\n", "else:\n", " wandb.login() # interactive prompt\n", "\n", "if HF_TOKEN:\n", " login(token=HF_TOKEN)\n", "\n", "os.environ[\"HF_HUB_MODEL_ID\"] = HF_HUB_MODEL_ID" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 4. Verify environment (smoke test)" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "from environment.env import ContextCorruptionEnv\n", "from environment.actions import ContextCorruptionAction, ActionType\n", "\n", "env = ContextCorruptionEnv(difficulty=2)\n", "obs = env.reset()\n", "assert len(obs.documents) == 8\n", "obs = env.step(ContextCorruptionAction(action_type=ActionType.submit_answer, answer=\"test\", confidence=0.5))\n", "assert obs.done and obs.reward is not None\n", "print(f\"✅ Smoke test passed | reward: {obs.reward:.4f}\")\n", "print(f\" Question: {env.state.question}\")" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 5. Preview training dataset" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "import sys\n", "sys.path.insert(0, \".\")\n", "from training.train_grpo import build_dataset, SYSTEM_PROMPT\n", "\n", "sample_ds = build_dataset(n_episodes=5, seed=0)\n", "sample = sample_ds[0]\n", "print(\"System:\", sample[\"messages\"][0][\"content\"][:200], \"...\")\n", "print(\"\\nUser message (first 400 chars):\", sample[\"messages\"][1][\"content\"][:400], \"...\")\n", "print(\"\\nGround truth:\", sample[\"ground_truth\"])\n", "print(\"Corrupt doc IDs:\", sample[\"corrupt_ids\"])" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 6. Run GRPO training\n", "\n", "Expected time on A100: ~45–60 min for 3 epochs over 500 episodes." ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "from training.train_grpo import main\n", "main()" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 7. View training curves" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "from IPython.display import Image, display\n", "\n", "display(Image(\"assets/reward_curve.png\"))\n", "display(Image(\"assets/loss_curve.png\"))" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 8. Evaluate trained model vs baseline" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "import json, torch, re\n", "from unsloth import FastLanguageModel\n", "from training.train_grpo import (\n", " MODEL_NAME, MAX_SEQ_LENGTH, OUTPUT_DIR,\n", " build_dataset, SYSTEM_PROMPT, _parse_completion\n", ")\n", "from environment.reward import compute_reward\n", "\n", "model, tokenizer = FastLanguageModel.from_pretrained(\n", " model_name=f\"{OUTPUT_DIR}-final\",\n", " max_seq_length=MAX_SEQ_LENGTH,\n", " load_in_4bit=True,\n", ")\n", "FastLanguageModel.for_inference(model)\n", "\n", "eval_ds = build_dataset(n_episodes=50, seed=999)\n", "rewards = []\n", "\n", "for row in eval_ds:\n", " prompt = tokenizer.apply_chat_template(\n", " row[\"messages\"], tokenize=False, add_generation_prompt=True\n", " )\n", " inputs = tokenizer(prompt, return_tensors=\"pt\").to(\"cuda\")\n", " with torch.no_grad():\n", " out = model.generate(**inputs, max_new_tokens=256, temperature=0.1, do_sample=True)\n", " completion = tokenizer.decode(out[0][inputs[\"input_ids\"].shape[1]:], skip_special_tokens=True)\n", " parsed = _parse_completion(completion)\n", " if parsed:\n", " reward, _ = compute_reward(\n", " parsed.get(\"answer\", \"\"), row[\"ground_truth\"],\n", " [int(x) for x in parsed.get(\"suspicious_docs\", [])],\n", " row[\"corrupt_ids\"], float(parsed.get(\"confidence\", 0.5)),\n", " budget_used=1, max_budget=12\n", " )\n", " else:\n", " reward = 0.0\n", " rewards.append(reward)\n", "\n", "avg = sum(rewards) / len(rewards)\n", "print(f\"\\n{'='*50}\")\n", "print(f\"Trained model avg reward : {avg:.4f}\")\n", "print(f\"Random baseline avg : 0.1302\")\n", "print(f\"Improvement : {avg - 0.1302:+.4f}\")\n", "print(f\"{'='*50}\")" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 9. Commit plots and results" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "trained_avg = avg # from cell above\n", "\n", "results = {\n", " \"baseline_avg_reward\": 0.1302,\n", " \"trained_avg_reward\": round(trained_avg, 4),\n", " \"improvement\": round(trained_avg - 0.1302, 4),\n", " \"n_eval_episodes\": 50,\n", " \"model\": \"Qwen2-1.5B-Instruct + LoRA r=16 GRPO\",\n", "}\n", "with open(\"eval/trained_results.json\", \"w\") as f:\n", " json.dump(results, f, indent=2)\n", "\n", "!git config user.email \"colab@training\"\n", "!git config user.name \"Colab Training Run\"\n", "!git add assets/reward_curve.png assets/loss_curve.png eval/trained_results.json\n", "!git commit -m \"results: add training curves and eval results\"\n", "!git push origin main\n", "print(\"Done — plots and results committed.\")" ] } ], "metadata": { "accelerator": "GPU", "colab": { "gpuType": "A100", "name": "ContextCorruption_GRPO.ipynb", "provenance": [] }, "kernelspec": { "display_name": "Python 3", "language": "python", "name": "python3" }, "language_info": { "name": "python", "version": "3.11.0" } }, "nbformat": 4, "nbformat_minor": 4 }