ms-swift评测插件拓展:如何自定义奖励函数?

在大模型强化学习训练中,奖励函数(Reward Function)是决定模型行为走向的“指挥棒”。它不只影响最终效果,更直接关系到对齐质量、安全边界与任务适配性。ms-swift 作为当前最活跃的大模型微调基础设施之一,其核心优势之一正是高度可插拔的评测与奖励体系——尤其在 GRPO、DPO、KTO 等人类对齐算法中,奖励函数不再是一个黑箱模块,而是可自由定义、热替换、即插即用的工程组件。

本文不讲抽象理论,不堆参数公式,而是以真实可运行的代码+分步调试逻辑+避坑经验,带你从零实现一个自定义奖励函数,并集成进 ms-swift 的评测与 RLHF 流程中。无论你是刚接触强化学习的新手,还是已在业务中落地 RLHF 的工程师,都能从中获得可立即复用的实践路径。


1. 为什么需要自定义奖励函数?

默认奖励函数(如基于 RM 模型打分、规则匹配关键词、长度惩罚等)虽开箱即用,但在实际场景中常面临三类典型瓶颈:

  • 领域失配:通用 RM 在金融/医疗/法律等垂直领域打分不准,需注入领域知识;
  • 多目标冲突:既要回答准确,又要语言简洁,还要符合品牌语气,单一标量难刻画;
  • 动态反馈缺失:用户真实点击、停留、修正行为无法被静态 RM 捕获,需对接线上日志流。

ms-swift 的插件化设计,正是为解决这些问题而生。它将奖励计算解耦为独立 Python 类,支持:

  • 纯 Python 实现(无需编译、无 CUDA 依赖)
  • 支持同步/异步调用(适配 vLLM 推理引擎)
  • 可与本地模型、API 服务、数据库、向量库任意组合
  • 自动参与梯度计算(若需可微)或仅用于采样排序(如 DPO)

下面,我们以一个技术文档问答质量评估器为例,逐步构建一个可落地的自定义奖励函数。


2. 奖励函数基础结构解析

ms-swift 中所有可插拔奖励函数均需继承 RewardPlugin 抽象基类。其核心接口极简:

from swift.llm import RewardPlugin

class MyRewardPlugin(RewardPlugin):
    def __init__(self, config: dict):
        super().__init__(config)
        # 初始化资源:加载小模型、连接DB、预热缓存等

    def compute_reward(self, query: str, response: str, history: list = None) -> float:
        """单次打分:输入query+response,返回标量reward"""
        raise NotImplementedError

    def compute_batch_reward(self, batch: list[dict]) -> list[float]:
        """批量打分(可选):提升吞吐,vLLM 异步推理时推荐实现"""
        return [self.compute_reward(**item) for item in batch]

注意:compute_reward 是必重写方法;compute_batch_reward 为性能优化项,非必须,但强烈建议实现(尤其在 GRPO 多轮采样中,batch size 常达 32–64)。

该结构天然支持两类使用模式:

  • 评测阶段:swift eval 调用 compute_reward 对生成结果打分,生成 reward 分布报告;
  • 训练阶段:swift rlhf 在 GRPO/DPO 等算法中调用,参与 loss 计算或样本筛选。

3. 实战:构建一个“技术文档准确性+可读性”双维度奖励器

我们以企业内部 LLM 助手为背景:用户提问 Kubernetes 配置问题,模型生成 YAML 片段。理想奖励应同时衡量:

  • 准确性:YAML 是否语法合法?是否包含必需字段(如 apiVersion, kind)?
  • 可读性:是否含冗余注释?行数是否超过合理阈值?是否使用了已弃用字段?

3.1 步骤一:创建奖励插件文件

新建文件 my_reward_plugin.py(路径任意,建议放在项目根目录或 plugins/ 子目录):

# my_reward_plugin.py
import re
import yaml
from typing import List, Dict, Any
from swift.llm import RewardPlugin

class TechDocRewardPlugin(RewardPlugin):
    def __init__(self, config: Dict[str, Any]):
        super().__init__(config)
        self.max_lines = config.get('max_lines', 50)
        self.required_keys = config.get('required_keys', ['apiVersion', 'kind', 'metadata'])
        self.deprecated_keys = config.get('deprecated_keys', ['spec.template.spec.restartPolicy'])

    def _is_valid_yaml(self, text: str) -> bool:
        try:
            yaml.safe_load(text)
            return True
        except (yaml.YAMLError, ValueError):
            return False

    def _count_lines(self, text: str) -> int:
        return len(text.strip().split('\n'))

    def _has_required_keys(self, text: str) -> bool:
        try:
            data = yaml.safe_load(text)
            if not isinstance(data, dict):
                return False
            return all(key in data for key in self.required_keys)
        except Exception:
            return False

    def _has_deprecated_keys(self, text: str) -> bool:
        for dep_key in self.deprecated_keys:
            if dep_key in text:
                return True
        return False

    def compute_reward(self, query: str, response: str, history: List[Dict] = None) -> float:
        score = 0.0

        # 准确性子项(权重 0.7)
        if self._is_valid_yaml(response):
            score += 0.4
        if self._has_required_keys(response):
            score += 0.3

        # 可读性子项(权重 0.3)
        lines = self._count_lines(response)
        if lines <= self.max_lines:
            score += 0.2
        if not self._has_deprecated_keys(response):
            score += 0.1

        # 惩罚项:含明显错误提示词
        if re.search(r'(error|invalid|not supported|deprecated)', response.lower()):
            score = max(0.0, score - 0.3)

        return round(score, 3)

    def compute_batch_reward(self, batch: List[Dict]) -> List[float]:
        return [self.compute_reward(**item) for item in batch]

该插件特点:

  • 完全纯 Python,无外部模型依赖,启动快、部署轻;
  • 支持配置化(max_lines, required_keys),便于 A/B 测试;
  • 包含基础容错(YAML 解析异常捕获);
  • 打分范围明确(0.0–1.0),便于后续归一化或加权。

4. 如何在评测中使用该奖励函数?

ms-swift 的 swift eval 命令原生支持通过 --reward_plugin 参数加载自定义插件。

4.1 准备评测数据集

创建 tech_qa_eval.jsonl(General-QA 格式):

{"query": "请给出一个 Deployment 的最小可用 YAML 示例", "response": "apiVersion: apps/v1\nkind: Deployment\nmetadata:\n  name: nginx-deploy"}
{"query": "如何设置 Pod 重启策略?", "response": "spec:\n  template:\n    spec:\n      restartPolicy: Always"}
{"query": "列出所有命名空间下的 Pod", "response": "kubectl get pods -A"}

注意:response 字段为模型生成内容,query 为原始输入。history 字段留空即可。

4.2 运行带自定义奖励的评测

CUDA_VISIBLE_DEVICES=0 \
swift eval \
  --model Qwen/Qwen2.5-7B-Instruct \
  --infer_backend pt \
  --eval_dataset no \
  --custom_eval_config tech_qa_eval.jsonl \
  --reward_plugin my_reward_plugin.TechDocRewardPlugin \
  --reward_plugin_config '{"max_lines": 40, "required_keys": ["apiVersion", "kind"]}' \
  --eval_output_dir eval_results_tech \
  --eval_limit 3

关键参数说明:

  • --reward_plugin: 模块路径 + 类名(格式:package.module.ClassName),本例中 my_reward_plugin.py 与命令同目录,故为 my_reward_plugin.TechDocRewardPlugin
  • --reward_plugin_config: JSON 字符串,传入 __init__ 的 config 参数
  • --custom_eval_config: 指向你的 .jsonl 文件(注意:此处非配置文件,而是数据文件,ms-swift 会自动识别 General-QA 格式)

执行后,你将在 eval_results_tech/ 下看到:

  • reward_scores.json: 每条样本的 reward 值列表
  • reward_summary.json: 平均分、标准差、分布直方图统计
  • detailed_results.jsonl: 每条 query-response-reward 的完整记录

示例输出节选:

[
  {"query": "请给出一个 Deployment 的最小可用 YAML 示例", "response": "apiVersion: apps/v1\nkind: Deployment\nmetadata:\n  name: nginx-deploy", "reward": 0.9},
  {"query": "如何设置 Pod 重启策略?", "response": "spec:\n  template:\n    spec:\n      restartPolicy: Always", "reward": 0.7},
  {"query": "列出所有命名空间下的 Pod", "response": "kubectl get pods -A", "reward": 0.0}
]

观察:第三条因 response 不是 YAML,reward 为 0 —— 这正是我们期望的“硬过滤”能力。


5. 如何在 GRPO 训练中集成该奖励函数?

GRPO(Generalized Reinforcement Learning with Policy Optimization)是 ms-swift 主推的异步强化学习框架,其核心流程为:

  1. 主模型(policy)生成多个 response(如 8 个);
  2. 奖励模型(reward plugin)对每个 response 打分;
  3. 根据 reward 排序,选择 top-k 作为正样本,bottom-k 为负样本;
  4. 构建 GRPO loss,更新 policy。

要启用自定义 reward 插件,只需在 swift rlhf 命令中添加相同参数:

CUDA_VISIBLE_DEVICES=0,1 NPROC_PER_NODE=2 \
swift rlhf \
  --rlhf_type grpo \
  --model Qwen/Qwen2.5-7B-Instruct \
  --train_type lora \
  --dataset AI-ModelScope/alpaca-gpt4-data-zh#1000 \
  --output_dir grpo_tech_finetune \
  --reward_plugin my_reward_plugin.TechDocRewardPlugin \
  --reward_plugin_config '{"max_lines": 45}' \
  --grpo_topk 4 \
  --grpo_bottomk 2 \
  --num_train_epochs 1 \
  --per_device_train_batch_size 2 \
  --gradient_accumulation_steps 8

关键点:

  • --grpo_topk / --grpo_bottomk 控制采样策略,数值需小于 --num_return_sequences(默认为 8);
  • reward plugin 会在每个 step 的 rollout 阶段被调用,自动批处理(调用 compute_batch_reward);
  • 若未实现 compute_batch_reward,则退化为逐条调用 compute_reward,性能下降但功能不变。

🧪 调试技巧:在 compute_reward 开头加入 print(f"[DEBUG] query: {query[:30]}... reward: {score}"),配合 --logging_steps 1 快速验证逻辑。


6. 高级技巧:组合多个奖励源

单一 reward 插件能力有限。ms-swift 支持奖励函数链式组合,通过 CompositeRewardPlugin 实现加权融合:

# composite_reward.py
from swift.llm import CompositeRewardPlugin, RewardPlugin

class CompositeTechReward(CompositeRewardPlugin):
    def __init__(self, config: dict):
        plugins_config = [
            {'plugin': 'my_reward_plugin.TechDocRewardPlugin', 'weight': 0.6, 'config': {'max_lines': 40}},
            {'plugin': 'swift.reward.rm_reward.RMRewardPlugin', 'weight': 0.3, 'config': {'model_id': 'Qwen/Qwen2.5-RM-1.5B'}},
            {'plugin': 'swift.reward.length_reward.LengthRewardPlugin', 'weight': 0.1}
        ]
        super().__init__(plugins_config)

使用方式完全一致:

--reward_plugin composite_reward.CompositeTechReward

这种设计让团队可并行开发:

  • 基础规则层(你写的 YAML 检查器)
  • 模型层(RM 打分)
  • 统计层(长度、重复率、困惑度)

各司其职,灵活组装。


7. 常见问题与避坑指南

问题现象根本原因解决方案
ModuleNotFoundError: No module named 'my_reward_plugin'Python 路径未包含插件所在目录启动前执行 export PYTHONPATH=$(pwd):$PYTHONPATH,或把插件放入 site-packages
reward 全为 0 或 NaNcompute_reward 抛出未捕获异常在 compute_reward 中包裹 try...except,返回默认分(如 0.1)并打印 error
GRPO 训练卡在 rollout 阶段compute_batch_reward 返回长度与 batch 不一致严格校验返回 list 长度等于 len(batch),建议用 assert len(ret) == len(batch)
reward 分布过于集中(如全 0.9)规则过宽松或阈值不合理在 eval 后查看 reward_summary.json,调整 max_lines、required_keys 等配置,或增加惩罚项
与 vLLM 异步推理不兼容插件内含阻塞 IO(如 requests.get)改用 httpx.AsyncClient + async def compute_batch_reward,ms-swift 自动识别异步方法

最佳实践:始终先用 swift eval 单独验证 reward 插件逻辑,再投入训练。一次验证胜过十次 debug。


8. 总结:掌握奖励函数,就是掌握对齐主动权

在 ms-swift 生态中,自定义奖励函数不是高阶技巧,而是基础工程能力。它意味着:

  • 你不再依赖“通用 RM”的模糊打分,而是用业务语言定义什么是好答案;
  • 你可以将产品指标(如用户停留时长、客服转人工率)映射为 reward,实现端到端优化;
  • 你能快速迭代对齐策略,今天加安全过滤,明天加风格约束,后天接入 A/B 测试平台。

本文所展示的 YAML 检查器,只是一个起点。你可以轻松扩展为:

  • 代码生成器:用 pyflakes 检查 Python 语法 + black 格式校验;
  • 财报分析助手:调用 pandas 解析表格文本 + 验证数字一致性;
  • 🧾 合同审查模型:用正则匹配关键条款 + 调用嵌入模型计算语义偏离度。

真正的 AI 工程化,始于对 reward 的掌控。而 ms-swift,为你提供了最平滑的落地路径。

---

> **获取更多AI镜像**
>
> 想探索更多AI镜像和应用场景?访问 [CSDN星图镜像广场](https://ai.csdn.net/?utm_source=mirror_blog_end),提供丰富的预置镜像,覆盖大模型推理、图像生成、视频生成、模型微调等多个领域,支持一键部署。
Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐