4b0c2077创建于 2025年3月26日历史提交
{
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# 奖励函数与数据准备\n",
    "\n",
    "本节详细介绍 Wordle 的奖励函数设计和训练数据准备流程。\n",
    "\n",
    "---\n",
    "\n",
    "## 1. 奖励函数\n",
    "\n",
    "Wordle 的奖励由 `wordle_reward.py` 中的 `compute_score` 函数计算,包含四个组件:\n",
    "\n",
    "| 组件 | 分值范围 | 说明 |\n",
    "|------|---------|------|\n",
    "| `correct_answer` | 0 或 1.0 | 猜中秘密单词 |\n",
    "| `partial_answer` | 0 - 0.8 | 部分匹配(0.2 * 绿色 + 0.1 * 黄色)|\n",
    "| `length_bonus` | 0 - 1.0 | 步数越少奖励越高(1 / 猜测次数)|\n",
    "| `format_reward` | 0 - 0.2 | 正确使用 `<guess>` 标签格式(权重 0.2)|\n",
    "\n",
    "### 奖励计算示例"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "```text\n",
    "秘密单词: crane\n",
    "\n",
    "Rollout A: 2 轮猜中\n",
    "  correct_answer = 1.0\n",
    "  partial_answer = 0.0  (猜中后不算 partial)\n",
    "  length_bonus   = 1/2 = 0.5\n",
    "  format_reward  = 1.0 * 0.2 = 0.2\n",
    "  总奖励 = 1.0 + 0.0 + 0.5 + 0.2 = 1.7\n",
    "\n",
    "Rollout B: 6 轮未猜中,最后一轮 3 绿 1 黄\n",
    "  correct_answer = 0.0\n",
    "  partial_answer = 0.2*3 + 0.1*1 = 0.7\n",
    "  length_bonus   = 0.0  (未猜中)\n",
    "  format_reward  = 1.0 * 0.2 = 0.2\n",
    "  总奖励 = 0.0 + 0.7 + 0.0 + 0.2 = 0.9\n",
    "```"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### 奖励函数返回值\n",
    "\n",
    "`compute_score` 返回一个字典,包含总分和各组件分数:"
   ]
  },
  {
   "cell_type": "code",
   "metadata": {},
   "source": [
    "# compute_score 返回值\n",
    "{\n",
    "    'score': 1.7,         # 总分(用于训练)\n",
    "    'correct': 1.0,       # 猜中奖励\n",
    "    'partial': 0.0,       # 部分匹配\n",
    "    'length_bonus': 0.5,  # 步数奖励\n",
    "    'format': 1.0,        # 格式正确率\n",
    "    'num_guesses': 2      # 猜测次数\n",
    "}\n"
   ],
   "outputs": [],
   "execution_count": null
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "验证时会分别记录每个组件到 tensorboard,方便分析模型的薄弱环节。\n",
    "\n",
    "---\n",
    "\n",
    "## 2. 数据准备\n",
    "\n",
    "训练数据由 `prepare_data.py` 生成,词表来源于 TextArena Wordle-v0。\n",
    "\n",
    "### 数据格式\n",
    "\n",
    "每条数据包含以下字段:"
   ]
  },
  {
   "cell_type": "code",
   "metadata": {},
   "source": [
    "# wordle_train.parquet 数据格式\n",
    "{\n",
    "    'prompt': [\n",
    "        {'role': 'system', 'content': 'You are a competitive game player...'},\n",
    "        {'role': 'user', 'content': '[GAME] You are Playing Wordle...'}\n",
    "    ],\n",
    "    'raw_prompt': [...],      # 同 prompt\n",
    "    'answer': 'crane',         # 秘密单词\n",
    "    'index': 'wordle_train_crane',\n",
    "    'data_source': 'wordle',\n",
    "    'reward_model': {'ground_truth': 'crane'}\n",
    "}\n"
   ],
   "outputs": [],
   "execution_count": null
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### 生成命令"
   ]
  },
  {
   "cell_type": "code",
   "metadata": {},
   "source": [
    "%%bash\n",
    "cd ~/rl-workspace/verl\n",
    "python3 prepare_data.py \\\n",
    "    --num_train 2000 --num_test 20 \\\n",
    "    --output_dir data\n"
   ],
   "execution_count": null,
   "outputs": []
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "生成 `data/wordle_train.parquet`(2000 条)和 `data/wordle_test.parquet`(20 条)。\n",
    "\n",
    "### 词表来源\n",
    "\n",
    "词表来自 TextArena Wordle-v0,TextArena 使用 NLTK 的 `pos_tag` 过滤出 5 字母名词作为有效词表。因此需要下载 NLTK 数据:\n",
    "- `words`:NLTK 英语词表\n",
    "- `averaged_perceptron_tagger_eng`:词性标注器\n",
    "\n",
    "---\n",
    "\n",
    "## 课后练习"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "1. (判断题)Wordle 奖励函数中,猜中后不再计算 partial_answer 分数。\n",
    "\n",
    "2. (判断题)length_bonus 的值与猜测次数成反比,猜得越快奖励越高。\n",
    "\n",
    "3. (判断题)format_reward 的权重是 1.0,与 correct_answer 相同。\n",
    "\n",
    "4. (单选题)一个 2 轮猜中的 rollout,其 length_bonus 是多少?\n",
    "    A. 0.2\n",
    "    B. 0.5\n",
    "    C. 1.0\n",
    "    D. 2.0\n",
    "\n",
    "5. (单选题)训练数据的词表来源于哪里?\n",
    "    A. Hugging Face 数据集\n",
    "    B. TextArena Wordle-v0\n",
    "    C. 手动标注\n",
    "    D. Wordle 官方网站\n",
    "\n",
    "6. (多选题)Wordle 奖励函数包含以下哪些组件?\n",
    "    A. correct_answer\n",
    "    B. partial_answer\n",
    "    C. length_bonus\n",
    "    D. format_reward"
   ]
  },
  {
   "cell_type": "code",
   "metadata": {},
   "source": [
    "!cat ./answer/03.03_answer.txt"
   ],
   "execution_count": null,
   "outputs": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "Python 3",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "name": "python",
   "version": "3.11"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 4
}