{ "cells": [ { "cell_type": "markdown", "metadata": {}, "source": [ "# 🚛 Driver Recruit Environment — RL Training with TRL\n", "\n", "Train a 3B LLM to recruit truck drivers using REINFORCE with TRL.\n", "\n", "The model learns to choose the right screening topics to ask drivers,\n", "then the auto-pilot handles CRM updates, approval, and hiring.\n", "\n", "**Environment**: [OpenEnv 0.2.1](https://github.com/meta-pytorch/OpenEnv) deployed on HF Spaces\n", "\n", "**Model**: Qwen/Qwen2.5-3B-Instruct\n", "\n", "**Algorithm**: REINFORCE with batch-level advantage normalization" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 1. Install Dependencies" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "!pip install -q openenv-core[core]==0.2.1 trl transformers torch accelerate" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 2. Connect to the Environment\n", "\n", "The recruiting environment is deployed on HF Spaces. Replace the URL below with your Space URL." ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "import json\n", "import random\n", "import re\n", "\n", "import torch\n", "import torch.nn.functional as F\n", "from transformers import AutoTokenizer, AutoModelForCausalLM\n", "\n", "# --- Connect to environment ---\n", "# Replace with your HF Space URL\n", "ENV_URL = \"https://YOUR-USERNAME-recruitopenenv.hf.space\" # <-- CHANGE THIS\n", "\n", "from openenv.client import EnvClient\n", "\n", "# Quick test: reset and check the env is alive\n", "import requests\n", "resp = requests.post(f\"{ENV_URL}/reset\", json={\"seed\": 42})\n", "data = resp.json()\n", "print(f\"Driver: {data['observation']['driver_name']}\")\n", "print(f\"Stage: {data['observation']['stage']}\")\n", "print(\"Environment connected!\")" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 3. Environment Helper Functions" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "def env_reset(seed=None):\n", " \"\"\"Reset environment via HTTP.\"\"\"\n", " payload = {\"seed\": seed} if seed else {}\n", " resp = requests.post(f\"{ENV_URL}/reset\", json=payload)\n", " return resp.json()\n", "\n", "def env_step(tool, action, **kwargs):\n", " \"\"\"Step environment via HTTP.\"\"\"\n", " payload = {\"tool\": tool, \"action\": action, **kwargs}\n", " resp = requests.post(f\"{ENV_URL}/step\", json=payload)\n", " return resp.json()\n", "\n", "# --- Topic-based auto-pilot ---\n", "SCREENING_TOPICS = [\n", " \"experience\", \"home_time\", \"pay\", \"equipment\", \"route\",\n", " \"deal_breakers\", \"availability\", \"violations\", \"medical_card\", \"references\",\n", "]\n", "\n", "SYSTEM_PROMPT = \"\"\"You are a truck driver recruiter screening a candidate. Choose the next topic to discuss.\n", "\n", "Topics for first contact: greeting (text), call (phone)\n", "Screening topics: experience, home_time, pay, equipment, route, deal_breakers, availability, violations, medical_card, references\n", "Say \"done\" when you have enough info to proceed with hiring.\n", "\n", "Respond with ONLY the topic name, nothing else.\"\"\"\n", "\n", "ALL_TOPICS = [\"greeting\", \"call\"] + SCREENING_TOPICS + [\"done\"]\n", "\n", "def parse_topic(text):\n", " \"\"\"Extract topic name from model output.\"\"\"\n", " text = text.strip().lower().replace('\"', '').replace(\"'\", \"\")\n", " text = text.split(\"\\n\")[0].strip().split(\".\")[0].strip()\n", " for topic in ALL_TOPICS:\n", " if topic in text or topic.replace(\"_\", \" \") in text:\n", " return topic\n", " if \"deal\" in text: return \"deal_breakers\"\n", " if \"home\" in text: return \"home_time\"\n", " if \"medical\" in text: return \"medical_card\"\n", " return \"done\"\n", "\n", "def build_prompt(obs, asked):\n", " \"\"\"Build prompt showing state and available topics.\"\"\"\n", " parts = [f\"Driver: {obs['driver_name']}\"]\n", " if obs.get('jobs_summary'):\n", " parts.append(f\"Jobs:\\n{obs['jobs_summary']}\")\n", " if obs.get('discovered_info'):\n", " parts.append(f\"Discovered:\\n{obs['discovered_info']}\")\n", " parts.append(f\"Stage: {obs['stage']}\")\n", " if asked:\n", " parts.append(f\"Already asked: {', '.join(asked)}\")\n", " available = [t for t in ALL_TOPICS if t not in asked]\n", " parts.append(f\"Available: {', '.join(available)}\")\n", " return \"\\n\".join(parts)\n", "\n", "print(\"Helpers loaded!\")" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 4. Run a Demo Episode\n", "\n", "Watch the auto-pilot run a full recruiting episode." ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "def run_demo_episode(seed=42):\n", " \"\"\"Run one full episode with scripted topic choices.\"\"\"\n", " state = env_reset(seed=seed)\n", " obs = state[\"observation\"]\n", " total_reward = 0.0\n", " print(f\"=== Driver: {obs['driver_name']} ===\")\n", "\n", " # Read CRM\n", " state = env_step(\"crm\", \"read_candidate\")\n", " total_reward += state[\"reward\"]\n", " obs = state[\"observation\"]\n", " print(f\"\\nJobs available:\\n{obs['jobs_summary'][:200]}...\")\n", "\n", " # Greet\n", " state = env_step(\"messaging\", \"send_message\", topic=\"greeting\")\n", " total_reward += state[\"reward\"]\n", " print(f\"\\nGreeting reward: {state['reward']}\")\n", "\n", " state = env_step(\"messaging\", \"read_reply\")\n", " total_reward += state[\"reward\"]\n", " obs = state[\"observation\"]\n", "\n", " state = env_step(\"crm\", \"update_stage\", stage=\"contacted\")\n", " total_reward += state[\"reward\"]\n", "\n", " # Screen\n", " for topic in [\"experience\", \"deal_breakers\", \"pay\", \"home_time\"]:\n", " if state.get(\"done\"): break\n", " state = env_step(\"messaging\", \"send_message\", topic=topic)\n", " total_reward += state[\"reward\"]\n", " state = env_step(\"messaging\", \"read_reply\")\n", " total_reward += state[\"reward\"]\n", " obs = state[\"observation\"]\n", " print(f\" {topic}: reward={state['reward']:.1f}\")\n", "\n", " print(f\"\\nDiscovered:\\n{obs.get('discovered_info', 'none')[:300]}\")\n", "\n", " # Approval + hire\n", " state = env_step(\"crm\", \"update_stage\", stage=\"interested\")\n", " total_reward += state[\"reward\"]\n", " state = env_step(\"approval\", \"request_approval\", job_id=0)\n", " total_reward += state[\"reward\"]\n", " state = env_step(\"workflow\", \"wait\")\n", " total_reward += state[\"reward\"]\n", " state = env_step(\"approval\", \"check_approval\")\n", " total_reward += state[\"reward\"]\n", " state = env_step(\"crm\", \"update_stage\", stage=\"approval_pending\")\n", " total_reward += state[\"reward\"]\n", " state = env_step(\"messaging\", \"send_message\", topic=\"offer\", job_id=0)\n", " total_reward += state[\"reward\"]\n", " state = env_step(\"messaging\", \"read_reply\")\n", " total_reward += state[\"reward\"]\n", " state = env_step(\"crm\", \"update_stage\", stage=\"offer_sent\")\n", " total_reward += state[\"reward\"]\n", " state = env_step(\"crm\", \"update_stage\", stage=\"hired\")\n", " total_reward += state[\"reward\"]\n", "\n", " obs = state[\"observation\"]\n", " print(f\"\\nFinal stage: {obs['stage']}\")\n", " print(f\"Total reward: {total_reward:.1f}\")\n", " print(f\"Done: {state.get('done')}\")\n", " return total_reward\n", "\n", "run_demo_episode()" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 5. Load Model" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "MODEL_NAME = \"Qwen/Qwen2.5-3B-Instruct\"\n", "TEMPERATURE = 1.5\n", "MAX_NEW_TOKENS = 32\n", "MAX_TOPICS = 8\n", "\n", "tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)\n", "if tokenizer.pad_token_id is None:\n", " tokenizer.pad_token_id = tokenizer.eos_token_id\n", "\n", "model = AutoModelForCausalLM.from_pretrained(\n", " MODEL_NAME,\n", " torch_dtype=torch.bfloat16,\n", " device_map=\"auto\",\n", ")\n", "model.gradient_checkpointing_enable()\n", "\n", "optimizer = torch.optim.AdamW(model.parameters(), lr=5e-6)\n", "device = next(model.parameters()).device\n", "print(f\"Model loaded on {device}\")" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 6. Training Loop — REINFORCE with Auto-Pilot\n", "\n", "The model only picks screening topics (1-5 tokens per decision).\n", "The auto-pilot handles CRM, stages, approval, and hiring.\n", "Rewards come from the full episode outcome." ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "def rollout_episode(model, tokenizer, device, seed=None):\n", " \"\"\"Run one auto-piloted episode. Model picks topics, wrapper does the rest.\"\"\"\n", " if seed is None:\n", " seed = random.randint(0, 2**31 - 1)\n", "\n", " state = env_reset(seed=seed)\n", " obs = state[\"observation\"]\n", " total_reward = 0.0\n", "\n", " # Auto: read CRM\n", " state = env_step(\"crm\", \"read_candidate\")\n", " total_reward += state[\"reward\"]\n", " obs = state[\"observation\"]\n", "\n", " if state.get(\"done\"):\n", " return None\n", "\n", " turn_data = []\n", " asked = []\n", " contacted = False\n", "\n", " for _ in range(MAX_TOPICS):\n", " if state.get(\"done\"):\n", " break\n", "\n", " # Build prompt\n", " obs_text = build_prompt(obs, asked)\n", " messages = [\n", " {\"role\": \"system\", \"content\": SYSTEM_PROMPT},\n", " {\"role\": \"user\", \"content\": obs_text},\n", " ]\n", " prompt = tokenizer.apply_chat_template(messages, add_generation_prompt=True, tokenize=False)\n", " input_ids = tokenizer.encode(prompt, return_tensors=\"pt\").to(device)\n", "\n", " # Model generates topic\n", " with torch.no_grad():\n", " output = model.generate(\n", " input_ids, max_new_tokens=MAX_NEW_TOKENS,\n", " do_sample=True, temperature=TEMPERATURE,\n", " pad_token_id=tokenizer.pad_token_id,\n", " )\n", " gen_ids = output[0, input_ids.shape[1]:].tolist()\n", " response = tokenizer.decode(gen_ids, skip_special_tokens=True)\n", " topic = parse_topic(response)\n", "\n", " turn_data.append({\n", " \"prompt_ids\": input_ids[0].tolist(),\n", " \"gen_ids\": gen_ids,\n", " \"topic\": topic,\n", " \"turn_reward\": 0.0,\n", " })\n", "\n", " if topic == \"done\":\n", " break\n", " if topic in asked:\n", " total_reward -= 0.5\n", " turn_data[-1][\"turn_reward\"] = -0.5\n", " asked.append(topic)\n", " continue\n", "\n", " asked.append(topic)\n", "\n", " # Auto: send_message + read_reply\n", " state = env_step(\"messaging\", \"send_message\", topic=topic)\n", " turn_reward = state[\"reward\"]\n", " total_reward += state[\"reward\"]\n", " obs = state[\"observation\"]\n", "\n", " if not state.get(\"done\"):\n", " state = env_step(\"messaging\", \"read_reply\")\n", " turn_reward += state[\"reward\"]\n", " total_reward += state[\"reward\"]\n", " obs = state[\"observation\"]\n", "\n", " # Auto: update stage after contact\n", " if topic in (\"greeting\", \"call\") and not contacted and not state.get(\"done\"):\n", " contacted = True\n", " state = env_step(\"crm\", \"update_stage\", stage=\"contacted\")\n", " turn_reward += state[\"reward\"]\n", " total_reward += state[\"reward\"]\n", " obs = state[\"observation\"]\n", "\n", " turn_data[-1][\"turn_reward\"] = turn_reward\n", "\n", " if not turn_data:\n", " return None\n", "\n", " # Auto: approval + offer + hire\n", " if not state.get(\"done\") and contacted:\n", " for action_spec in [\n", " (\"crm\", \"update_stage\", {\"stage\": \"interested\"}),\n", " (\"approval\", \"request_approval\", {\"job_id\": 0}),\n", " (\"workflow\", \"wait\", {}),\n", " (\"approval\", \"check_approval\", {}),\n", " (\"crm\", \"update_stage\", {\"stage\": \"approval_pending\"}),\n", " (\"messaging\", \"send_message\", {\"topic\": \"offer\", \"job_id\": 0}),\n", " (\"messaging\", \"read_reply\", {}),\n", " (\"crm\", \"update_stage\", {\"stage\": \"offer_sent\"}),\n", " (\"crm\", \"update_stage\", {\"stage\": \"hired\"}),\n", " ]:\n", " if state.get(\"done\"): break\n", " state = env_step(action_spec[0], action_spec[1], **action_spec[2])\n", " total_reward += state[\"reward\"]\n", "\n", " # Sample one turn for training\n", " t = random.randrange(len(turn_data))\n", " td = turn_data[t]\n", "\n", " return {\n", " \"prompt_ids\": td[\"prompt_ids\"],\n", " \"gen_ids\": td[\"gen_ids\"][:MAX_NEW_TOKENS],\n", " \"reward\": total_reward,\n", " \"stage\": obs.get(\"stage\", \"unknown\"),\n", " \"topic\": td[\"topic\"],\n", " \"num_topics\": len(asked),\n", " }\n", "\n", "# Quick test\n", "ep = rollout_episode(model, tokenizer, device)\n", "if ep:\n", " print(f\"Topic chosen: {ep['topic']}, Reward: {ep['reward']:.1f}, Stage: {ep['stage']}, Topics asked: {ep['num_topics']}\")" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "# --- REINFORCE Training Loop ---\n", "BATCH_SIZE = 4\n", "NUM_STEPS = 50\n", "\n", "print(f\"Training for {NUM_STEPS} steps, batch size {BATCH_SIZE}\")\n", "print(\"=\" * 60)\n", "\n", "history = {\"loss\": [], \"reward_mean\": [], \"reward_std\": [], \"grad_norm\": []}\n", "\n", "for step in range(1, NUM_STEPS + 1):\n", " # --- Rollout (no gradients) ---\n", " model.eval()\n", " episodes = []\n", " for i in range(BATCH_SIZE):\n", " ep = rollout_episode(model, tokenizer, device)\n", " if ep and ep[\"gen_ids\"]:\n", " episodes.append(ep)\n", "\n", " if len(episodes) < 2:\n", " print(f\"Step {step}: not enough episodes, skipping\")\n", " continue\n", "\n", " # --- Batch-level advantages ---\n", " rewards = [ep[\"reward\"] for ep in episodes]\n", " mean_r = sum(rewards) / len(rewards)\n", " std_r = max(torch.tensor(rewards).std().item(), 1e-4)\n", " advantages = [(r - mean_r) / std_r for r in rewards]\n", "\n", " # --- REINFORCE update ---\n", " model.train()\n", " optimizer.zero_grad()\n", " total_loss = 0.0\n", "\n", " for ep, adv in zip(episodes, advantages):\n", " input_ids = torch.tensor(\n", " [ep[\"prompt_ids\"] + ep[\"gen_ids\"]], device=device\n", " )\n", " prompt_len = len(ep[\"prompt_ids\"])\n", " comp_len = len(ep[\"gen_ids\"])\n", " if comp_len == 0:\n", " continue\n", "\n", " outputs = model(input_ids)\n", " logits = outputs.logits[0, prompt_len - 1 : prompt_len + comp_len - 1]\n", " targets = input_ids[0, prompt_len : prompt_len + comp_len]\n", " log_probs = F.log_softmax(logits, dim=-1)\n", " token_lps = log_probs.gather(1, targets.unsqueeze(1)).squeeze(1)\n", "\n", " loss = -(adv * token_lps.sum()) / len(episodes)\n", " loss.backward()\n", " total_loss += loss.item()\n", "\n", " grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0).item()\n", " optimizer.step()\n", "\n", " # --- Log ---\n", " history[\"loss\"].append(total_loss)\n", " history[\"reward_mean\"].append(mean_r)\n", " history[\"reward_std\"].append(std_r)\n", " history[\"grad_norm\"].append(grad_norm)\n", "\n", " topics = [ep[\"topic\"] for ep in episodes]\n", " print(f\"Step {step:3d} | loss={total_loss:+.3f} | reward={mean_r:+.1f}±{std_r:.1f} | \"\n", " f\"grad={grad_norm:.3f} | topics={topics}\")\n", "\n", "print(\"\\nTraining complete!\")" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 7. Plot Training Curves" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "import matplotlib.pyplot as plt\n", "\n", "fig, axes = plt.subplots(2, 2, figsize=(12, 8))\n", "fig.suptitle(\"Driver Recruit RL Training\", fontsize=14)\n", "\n", "axes[0, 0].plot(history[\"loss\"])\n", "axes[0, 0].set_title(\"Loss\")\n", "axes[0, 0].set_xlabel(\"Step\")\n", "\n", "axes[0, 1].plot(history[\"reward_mean\"])\n", "axes[0, 1].set_title(\"Mean Reward\")\n", "axes[0, 1].set_xlabel(\"Step\")\n", "\n", "axes[1, 0].plot(history[\"reward_std\"])\n", "axes[1, 0].set_title(\"Reward Std\")\n", "axes[1, 0].set_xlabel(\"Step\")\n", "\n", "axes[1, 1].plot(history[\"grad_norm\"])\n", "axes[1, 1].set_title(\"Gradient Norm\")\n", "axes[1, 1].set_xlabel(\"Step\")\n", "\n", "plt.tight_layout()\n", "plt.show()" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 8. Test the Trained Model" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "print(\"=== Testing trained model ===\")\n", "model.eval()\n", "test_rewards = []\n", "for i in range(5):\n", " ep = rollout_episode(model, tokenizer, device)\n", " if ep:\n", " test_rewards.append(ep[\"reward\"])\n", " print(f\" Episode {i+1}: reward={ep['reward']:.1f}, stage={ep['stage']}, \"\n", " f\"topics={ep['num_topics']}, chose={ep['topic']}\")\n", "\n", "if test_rewards:\n", " print(f\"\\nMean test reward: {sum(test_rewards)/len(test_rewards):.1f}\")" ] } ], "metadata": { "kernelspec": { "display_name": "Python 3", "language": "python", "name": "python3" }, "language_info": { "name": "python", "version": "3.10.0" }, "accelerator": "GPU", "colab": { "gpuType": "T4", "provenance": [] } }, "nbformat": 4, "nbformat_minor": 4 }