diff --git a/swift/infer_engine/protocol.py b/swift/infer_engine/protocol.py index 88bf48f518..8fd25ef483 100644 --- a/swift/infer_engine/protocol.py +++ b/swift/infer_engine/protocol.py @@ -156,6 +156,8 @@ class RolloutInferRequest(InferRequest): images: List[str] = field(default_factory=list) data_dict: Dict = field(default_factory=dict) uuid: Optional[str] = None + # 'auto' uses the training suffix policy when scoring completed responses. + add_eos: Optional[Union[bool, Literal['auto']]] = None def random_uuid() -> str: diff --git a/swift/ray/megatron/gkd_trainer.py b/swift/ray/megatron/gkd_trainer.py index 2d55c5fc74..44c25fca8c 100644 --- a/swift/ray/megatron/gkd_trainer.py +++ b/swift/ray/megatron/gkd_trainer.py @@ -325,6 +325,7 @@ def _fetch_teacher_from_replicas(self, gkd_samples: List[GKDSample], samples): teacher_encodeds.append(teacher_encoded) else: teacher_encodeds.append(None) + req.add_eos = 'auto' if s.add_eos is None else s.add_eos requests.append(req) request_config = RequestConfig(prompt_logprobs=topk, max_tokens=1, temperature=0.0) diff --git a/swift/rl_core/data.py b/swift/rl_core/data.py index 5da02bb973..5455fb22b9 100644 --- a/swift/rl_core/data.py +++ b/swift/rl_core/data.py @@ -125,7 +125,7 @@ def to_teacher_template_dict(self) -> Dict[str, Any]: d['chat_template_kwargs'] = chat_template_kwargs if self.response_token_ids: d['response_token_ids'] = self.response_token_ids - d['add_eos'] = False + d['add_eos'] = self.add_eos return d def _standard_fields(self) -> Dict[str, Any]: @@ -344,7 +344,8 @@ def to_device(self, device) -> 'GRPOBatch': @dataclass class GKDSample(OnPolicySample): - pass + # Dataset responses use the template's EOS policy; rollout outputs set this to False. + add_eos: Optional[bool] = None @dataclass diff --git a/swift/rlhf_trainers/gkd_helpers.py b/swift/rlhf_trainers/gkd_helpers.py index 7462700af8..c91c363e90 100644 --- a/swift/rlhf_trainers/gkd_helpers.py +++ b/swift/rlhf_trainers/gkd_helpers.py @@ -117,6 +117,8 @@ def build_teacher_requests(samples: List[OnPolicySample], template: Optional[Tem loss_mask, non_thinking_prefix_ids=prefix_ids) req.messages = messages + if isinstance(s, GKDSample): + req.add_eos = 'auto' if s.add_eos is None else s.add_eos requests.append(req) return requests diff --git a/swift/template/base.py b/swift/template/base.py index 2a0ce04f51..61e003599f 100644 --- a/swift/template/base.py +++ b/swift/template/base.py @@ -1443,8 +1443,8 @@ def _swift_encode(self, inputs: StdTemplateInputs): # self.is_training needed because we may want to continue generation from # the current response add_eos = inputs.extra_kwargs.get('add_eos') - if add_eos is None: - add_eos = (self.is_training + if add_eos is None or add_eos == 'auto': + add_eos = (self.is_training or add_eos == 'auto' or self.task_type != 'causal_lm') and not sep_token and not endswith_stop_words if add_eos: extra_context_list = template_meta.suffix