【Bug已解决】[AsyncGRPO] AsyncGRPOTrainer ignores provided processing_class when initializing AsyncRollou

【Bug已解决】[AsyncGRPO] AsyncGRPOTrainer ignores provided processing_class when initializing AsyncRollou

【Bug已解决】[AsyncGRPO] AsyncGRPOTrainer ignores provided processing_class when initializing AsyncRolloutWorker 解决方案

一、现象长什么样

AsyncGRPOTrainer做异步 RL,我们显式传了一个自定义processing_class(比如带特殊 chat template 的 processor / tokenizer),但训练出来模型的行为和预期不符,且 rollout 阶段生成的文本和训练阶段对不上。排查发现:AsyncRolloutWorker(负责用 vLLM 做生成的异步 rollout 工作进程)根本没有用到我们传的processing_class,而是自行初始化了一个默认的

现象:

trainer 用的 processing_class: 自定义 processor (特殊 chat template) AsyncRolloutWorker 用的: 默认 tokenizer (普通模板) -> 生成与训练 tokenization 不一致,reward 算错 / 行为偏移

更隐蔽的是:如果默认 processor 和自定义 processor 的词表或 chat template 不同,rollout 生成的 token id 和训练侧算 logps 时的 token id 对不上,GRPO 的 loss 直接建立在错误对齐上,训练静默失效。

现象特征:

  • 只有 AsyncGRPO 暴露(同步 GRPOTrainer 没有这个"worker 独立初始化"的问题);
  • 不报错,但生成/训练 tokenization 错位;
  • 用户传了processing_class却"好像没传一样"。

二、背景

AsyncGRPOTrainer的异步架构里,生成(rollout)和训练是分离的两个世界

  • 训练侧:在主进程/训练器里,用processing_class把样本编码成模型输入,算 logps、loss;
  • rollout 侧AsyncRolloutWorker(常是独立进程,甚至多进程池)用 vLLM 做自回归生成,它自己也需要一个 processor 来把 prompt 编码、把生成结果解码。

关键点:两侧的 processor 必须完全一致——同一个 chat template、同一个词表、同一个特殊 token 定义。否则 rollout 生成的 token 序列,在训练侧按"另一个 processor"解码/算 logps 时会错位,GRPO 的"生成-训练"闭环就建立在错误对齐上。

AsyncGRPOTrainer在初始化时会收processing_class参数,但把它只存给了训练侧,没传给AsyncRolloutWorker的初始化——worker 自己AutoProcessor.from_pretrained(model_name)拉了一个默认的。当用户提供了自定义 processor(特殊模板),worker 就用了错的那个。

三、根因

根因一句话:AsyncGRPOTrainer在初始化AsyncRolloutWorker时,没有把用户提供的processing_class透传进去,worker 自行用模型名加载了默认 processor,导致 rollout 侧与训练侧使用了不一致的 processor(chat template/词表/特殊 token 不同),生成与训练的 tokenization 错位,GRPO 闭环建立在错误对齐上,训练静默失效

具体:

  1. 透传缺失self.processing_class存了,但AsyncRolloutWorker(...)没收到它;
  2. worker 自加载默认:worker 用from_pretrained(model)拉默认 processor,忽略用户自定义;
  3. 两侧不一致:训练侧特殊模板,rollout 侧默认模板,token id 错位;
  4. 只在 AsyncGRPO 暴露:同步 GRPOTrainer 生成和训练在同一进程同一 processor,没这个问题;
  5. 静默失效:不报错,但 loss 建立在错位 token 上,reward 失真。

本质是"异步架构下,配置(processor)没有跨进程/跨组件透传"。

四、最小可运行复现

下面用纯 Python 模拟"worker 没收到 processing_class 用了默认":

class Processor: def __init__(self, template): self.template = template class AsyncRolloutWorker: def __init__(self, model_name, processing_class=None): if processing_class is None: # 旧行为:没收到就自加载默认 self.processor = Processor(template="default") else: self.processor = processing_class class AsyncGRPOTrainer: def __init__(self, model_name, processing_class=None): self.processing_class = processing_class # 旧实现:初始化 worker 时没传 processing_class self.worker = AsyncRolloutWorker(model_name) # ← 漏传 def check_consistency(self): train_tpl = getattr(self.processing_class, "template", "default") rollout_tpl = self.worker.processor.template return train_tpl == rollout_tpl def demo(): custom = Processor(template="special_chat") trainer = AsyncGRPOTrainer("my-model", processing_class=custom) print("训练侧模板:", custom.template) print("rollout(worker)模板:", trainer.worker.processor.template) print("两侧一致:", trainer.check_consistency(), " <- False 即错位") if __name__ == "__main__": demo()

输出:

训练侧模板: special_chat rollout(worker)模板: default 两侧一致: False <- False 即错位

rollout侧用了default而非用户传的special_chat,两侧不一致。复现了核心 bug:worker 没收到 processing_class。

五、解决方案(第一层):初始化 worker 时透传 processing_class

第一层最直接:把self.processing_class透传给AsyncRolloutWorker的初始化:

class AsyncRolloutWorker: def __init__(self, model_name, processing_class=None): # 收到就用用户的,没收到才默认 self.processor = processing_class if processing_class is not None \ else Processor(template="default") class AsyncGRPOTrainer: def __init__(self, model_name, processing_class=None): self.processing_class = processing_class # 修复:把 processing_class 透传给 worker self.worker = AsyncRolloutWorker(model_name, processing_class=self.processing_class) def check_consistency(self): train_tpl = getattr(self.processing_class, "template", "default") rollout_tpl = self.worker.processor.template return train_tpl == rollout_tpl def demo(): custom = Processor(template="special_chat") trainer = AsyncGRPOTrainer("my-model", processing_class=custom) print("修复后两侧一致:", trainer.check_consistency()) # True if __name__ == "__main__": demo()

核心是AsyncRolloutWorker(model_name, processing_class=self.processing_class)——用户传了就透传,worker 用同一份 processor,两侧一致。

六、解决方案(第二层):worker 内部用传入的 processor 而非重新加载

第一层修了透传,但要保证 worker 内部确实使用传入的 processor 做编解码,而不是"收了又去from_pretrained覆盖"。第二层在 worker 里把 processor 作为唯一真源:

class AsyncRolloutWorker: def __init__(self, model_name, processing_class=None): if processing_class is not None: self.processor = processing_class # 直接用传入的,不重载 else: self.processor = Processor(template="default") def encode_prompt(self, text): # 用 self.processor 编码,保证和训练侧同一模板 return f"[{self.processor.template}]{text}" def decode_completion(self, ids): # 同理用同一 processor 解码 return f"decode-by-{self.processor.template}" class AsyncGRPOTrainer: def __init__(self, model_name, processing_class=None): self.processing_class = processing_class self.worker = AsyncRolloutWorker(model_name, processing_class=self.processing_class) def rollout_and_train_consistent(self, prompt): encoded = self.worker.encode_prompt(prompt) # 训练侧也用 self.processing_class 编码,保证同源 train_encoded = f"[{self.processing_class.template}]{prompt}" if self.processing_class else prompt return encoded == train_encoded def demo(): custom = Processor(template="special_chat") t = AsyncGRPOTrainer("m", processing_class=custom) print("编解码同源:", t.rollout_and_train_consistent("hi")) if __name__ == "__main__": demo()
  • worker 收到processing_class不再from_pretrained重载,直接用传入实例;
  • 编码、解码都用self.processor,与训练侧self.processing_class同源;
  • 消除"收了又覆盖"的隐患。

七、解决方案(第三层):一致性断言 + 不变量测试

第三层加护栏,确保"训练侧与 rollout 侧 processor 完全一致",并锁进测试:

def assert_processor_consistent(trainer_processor, worker_processor): """断言两侧 processor 是同一实例或等价配置。""" if trainer_processor is None and worker_processor is None: return True if (trainer_processor is None) != (worker_processor is None): raise AssertionError("训练侧与 rollout 侧 processor 存在性不一致") # 比关键配置:模板 / 词表大小 t1 = getattr(trainer_processor, "template", None) t2 = getattr(worker_processor, "template", None) if t1 != t2: raise AssertionError(f"processor 模板不一致: 训练={t1} rollout={t2}") return True def test_processing_class_propagates(): custom = Processor(template="special_chat") t = AsyncGRPOTrainer("m", processing_class=custom) assert_processor_consistent(t.processing_class, t.worker.processor) print("OK: processing_class 已透传到 worker,两侧一致") def test_default_fallback_consistent(): # 不传 processing_class 时,两侧都用默认,仍一致 t = AsyncGRPOTrainer("m", processing_class=None) assert_processor_consistent(t.processing_class, t.worker.processor) print("OK: 默认路径两侧也一致") if __name__ == "__main__": test_processing_class_propagates() test_default_fallback_consistent()
  • assert_processor_consistent在 trainer 初始化后检查训练侧与 worker 侧 processor 的模板/词表一致,不一致立即断言失败;
  • 两个测试锁住"传自定义则透传一致"和"不传则默认也一致",任何把透传改回"worker 自加载默认"的改动被 CI 拦下。

八、落地建议

如果你在 AsyncGRPOTrainer 上发现生成/训练错位,建议:

  1. 确认 worker 是否收到 processing_class:没透传就加processing_class=self.processing_class
  2. worker 不重载:收到后直接用传入实例,别from_pretrained覆盖。
  3. 编解码同源:worker 编码/解码都用self.processor
  4. 加一致性断言:初始化后检查训练侧/worker 侧 processor 模板一致。
  5. 加测试:锁住"传则透传一致""不传则默认一致"。
  6. 多进程注意:若 worker 是独立进程,processor 需可序列化/可重建为等价实例。

九、排查清单

如果 AsyncGRPO 生成与训练错位,按顺序查:

  1. 确认处理 class 是否透传AsyncRolloutWorker(...)是否收到processing_class
  2. worker 是否自加载默认:收到后是否又from_pretrained覆盖。
  3. 看两侧模板/词表:训练侧self.processing_class与 workerself.processor是否一致。
  4. 编解码同源:worker 用self.processor做编解码。
  5. 加一致性断言:初始化后检查两侧 processor 一致。
  6. 加测试:锁住"透传一致""默认一致"。
  7. 多进程序列化:worker 跨进程时 processor 需可重建等价实例。

十、小结

AsyncGRPOTrainer忽略用户提供的processing_class,根因是trainer 初始化AsyncRolloutWorker时没有把processing_class透传进去,worker 自行用模型名加载了默认 processor,导致 rollout 侧与训练侧使用不一致的 processor(chat template/词表/特殊 token 不同),生成与训练的 tokenization 错位,GRPO 的"生成-训练"闭环建立在错误对齐上,训练静默失效。它只在异步架构暴露(同步 GRPOTrainer 同进程同 processor),且不报错,最难察觉。

修复分三层:第一层在AsyncRolloutWorker(model_name, processing_class=self.processing_class)透传,用户传了就给 worker;第二层让 worker 收到后直接用传入实例、不再from_pretrained重载,编解码都用self.processor,与训练侧同源;第三层加assert_processor_consistent初始化后检查两侧模板/词表一致,并加"传则透传一致""不传则默认一致"不变量测试。核心心法是:异步/多进程架构下,任何影响"两侧语义一致性"的配置(processor、chat template、词表)都必须显式跨组件透传且禁止单侧自加载默认——否则 rollout 与训练会悄悄用不同processor,让整个 RL 闭环建立在错位 token 上,训练不报错却全面失真