From e09408e64f66fedf5aa4330885d54038da641e04 Mon Sep 17 00:00:00 2001 From: taking-lying-flat <1615405@qq.com> Date: Fri, 2 Oct 2026 14:38:21 +0800 Subject: [PATCH 1/3] fix: preserve EOS supervision in off-policy GKD --- swift/ray/megatron/gkd_trainer.py | 1 + swift/rl_core/data.py | 5 +++-- swift/rlhf_trainers/gkd_helpers.py | 2 ++ swift/template/base.py | 5 +++-- 4 files changed, 9 insertions(+), 4 deletions(-) diff --git a/swift/ray/megatron/gkd_trainer.py b/swift/ray/megatron/gkd_trainer.py index 2d55c5fc74..072cd25650 100644 --- a/swift/ray/megatron/gkd_trainer.py +++ b/swift/ray/megatron/gkd_trainer.py @@ -319,6 +319,7 @@ def _fetch_teacher_from_replicas(self, gkd_samples: List[GKDSample], samples): teacher_encodeds = [] # teacher-side encoded (OPSD) or None (non-OPSD) for s, sample in zip(gkd_samples, samples): req = s.to_infer_request() + req.chat_template_kwargs = {**req.chat_template_kwargs, 'add_eos': s.add_eos} teacher_encoded = sample.get('teacher_encoded') if s.teacher_messages: req.messages = s.teacher_messages 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..bf87a0dbe2 100644 --- a/swift/rlhf_trainers/gkd_helpers.py +++ b/swift/rlhf_trainers/gkd_helpers.py @@ -100,6 +100,8 @@ def build_teacher_requests(samples: List[OnPolicySample], template: Optional[Tem request_sample = copy.copy(s) request_sample.images = s.teacher_images req = request_sample.to_infer_request() + if isinstance(s, GKDSample): + req.chat_template_kwargs = {**req.chat_template_kwargs, 'add_eos': s.add_eos} # OPSD: score the teacher on its privileged prompt instead of the student prompt. teacher_messages = getattr(s, 'teacher_messages', None) messages = teacher_messages if teacher_messages else req.messages diff --git a/swift/template/base.py b/swift/template/base.py index 2a0ce04f51..5e5fd9cc9a 100644 --- a/swift/template/base.py +++ b/swift/template/base.py @@ -1442,9 +1442,10 @@ def _swift_encode(self, inputs: StdTemplateInputs): if isinstance(stop_word, str)) # self.is_training needed because we may want to continue generation from # the current response - add_eos = inputs.extra_kwargs.get('add_eos') + add_eos = inputs.extra_kwargs.get('add_eos', inputs.chat_template_kwargs.get('add_eos')) if add_eos is None: - add_eos = (self.is_training + # Teacher scoring can request the automatic training EOS policy via template kwargs. + add_eos = (self.is_training or 'add_eos' in inputs.chat_template_kwargs or self.task_type != 'causal_lm') and not sep_token and not endswith_stop_words if add_eos: extra_context_list = template_meta.suffix From cccf842fff5c1120db1a00343dcc0cd8771bb161 Mon Sep 17 00:00:00 2001 From: taking-lying-flat <1615405@qq.com> Date: Fri, 2 Oct 2026 17:29:05 +0800 Subject: [PATCH 2/3] fix: resolve GKD teacher EOS without changing templates --- swift/infer_engine/protocol.py | 1 + swift/ray/megatron/gkd_trainer.py | 3 ++- swift/rlhf_trainers/gkd_helpers.py | 23 ++++++++++++++++++++--- swift/template/base.py | 5 ++--- 4 files changed, 25 insertions(+), 7 deletions(-) diff --git a/swift/infer_engine/protocol.py b/swift/infer_engine/protocol.py index 88bf48f518..4eb3873d40 100644 --- a/swift/infer_engine/protocol.py +++ b/swift/infer_engine/protocol.py @@ -156,6 +156,7 @@ class RolloutInferRequest(InferRequest): images: List[str] = field(default_factory=list) data_dict: Dict = field(default_factory=dict) uuid: Optional[str] = None + add_eos: Optional[bool] = None def random_uuid() -> str: diff --git a/swift/ray/megatron/gkd_trainer.py b/swift/ray/megatron/gkd_trainer.py index 072cd25650..8d57dcb30c 100644 --- a/swift/ray/megatron/gkd_trainer.py +++ b/swift/ray/megatron/gkd_trainer.py @@ -11,6 +11,7 @@ from swift.infer_engine.protocol import RequestConfig, RolloutOutput from swift.rl_core.data import GKDSample +from swift.rlhf_trainers.gkd_helpers import set_teacher_request_eos from swift.rlhf_trainers.gkd_loss import DataSource, TeacherOutput from swift.rlhf_trainers.utils import parse_prompt_logprobs from swift.rollout import MultiTurnScheduler, invoke_async_hook, multi_turns, run_multi_turn @@ -319,13 +320,13 @@ def _fetch_teacher_from_replicas(self, gkd_samples: List[GKDSample], samples): teacher_encodeds = [] # teacher-side encoded (OPSD) or None (non-OPSD) for s, sample in zip(gkd_samples, samples): req = s.to_infer_request() - req.chat_template_kwargs = {**req.chat_template_kwargs, 'add_eos': s.add_eos} teacher_encoded = sample.get('teacher_encoded') if s.teacher_messages: req.messages = s.teacher_messages teacher_encodeds.append(teacher_encoded) else: teacher_encodeds.append(None) + set_teacher_request_eos(req, s, self.template) requests.append(req) request_config = RequestConfig(prompt_logprobs=topk, max_tokens=1, temperature=0.0) diff --git a/swift/rlhf_trainers/gkd_helpers.py b/swift/rlhf_trainers/gkd_helpers.py index bf87a0dbe2..b57db81826 100644 --- a/swift/rlhf_trainers/gkd_helpers.py +++ b/swift/rlhf_trainers/gkd_helpers.py @@ -7,11 +7,12 @@ """ import copy import torch -from dataclasses import dataclass, field +from dataclasses import asdict, dataclass, field from typing import Any, Dict, List, Optional, Tuple from swift.rl_core.data import GKDSample, OnPolicySample from swift.template.base import Template +from swift.template.template_inputs import StdTemplateInputs from swift.utils import get_cu_seqlens_from_position_ids, get_logger, json_parse_to_dict from .gkd_loss import TeacherOutput from .utils import (assemble_teacher_topk_logprobs, encode_sample, get_response_prefix_ids, @@ -80,6 +81,22 @@ def encode_gkd_samples( return student_encoded_list, teacher_encoded_list, has_opsd +def set_teacher_request_eos(request, sample: GKDSample, template: Optional[Template] = None) -> None: + """Resolve the training suffix policy before sending a GKD teacher request.""" + request.add_eos = sample.add_eos + if request.add_eos is not None or template is None or template.template_backend != 'swift': + return + # Render on a copy to reuse model-specific stop/overlap rules without changing + # the shared inference template or running multimodal processors again. + template = copy.copy(template) + template.set_mode('train') + inputs = StdTemplateInputs.from_dict(asdict(request)) + template._swift_prepare_inputs(inputs) + _, _, answer_len = template._swift_encode(inputs) + # The final answer consists of its response plus any appended suffix contexts. + request.add_eos = answer_len > 1 + + def build_teacher_requests(samples: List[OnPolicySample], template: Optional[Template] = None) -> List[Any]: """Build teacher API requests from samples (GKD or GRPO/OPD-RL). @@ -100,8 +117,6 @@ def build_teacher_requests(samples: List[OnPolicySample], template: Optional[Tem request_sample = copy.copy(s) request_sample.images = s.teacher_images req = request_sample.to_infer_request() - if isinstance(s, GKDSample): - req.chat_template_kwargs = {**req.chat_template_kwargs, 'add_eos': s.add_eos} # OPSD: score the teacher on its privileged prompt instead of the student prompt. teacher_messages = getattr(s, 'teacher_messages', None) messages = teacher_messages if teacher_messages else req.messages @@ -119,6 +134,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): + set_teacher_request_eos(req, s, template) requests.append(req) return requests diff --git a/swift/template/base.py b/swift/template/base.py index 5e5fd9cc9a..2a0ce04f51 100644 --- a/swift/template/base.py +++ b/swift/template/base.py @@ -1442,10 +1442,9 @@ def _swift_encode(self, inputs: StdTemplateInputs): if isinstance(stop_word, str)) # self.is_training needed because we may want to continue generation from # the current response - add_eos = inputs.extra_kwargs.get('add_eos', inputs.chat_template_kwargs.get('add_eos')) + add_eos = inputs.extra_kwargs.get('add_eos') if add_eos is None: - # Teacher scoring can request the automatic training EOS policy via template kwargs. - add_eos = (self.is_training or 'add_eos' in inputs.chat_template_kwargs + add_eos = (self.is_training or self.task_type != 'causal_lm') and not sep_token and not endswith_stop_words if add_eos: extra_context_list = template_meta.suffix From 1dbb971d1c85a7835f8ab69f5dbc425416627287 Mon Sep 17 00:00:00 2001 From: taking-lying-flat <1615405@qq.com> Date: Fri, 2 Oct 2026 17:48:17 +0800 Subject: [PATCH 3/3] fix: use explicit auto EOS policy for teacher scoring --- swift/infer_engine/protocol.py | 3 ++- swift/ray/megatron/gkd_trainer.py | 3 +-- swift/rlhf_trainers/gkd_helpers.py | 21 ++------------------- swift/template/base.py | 4 ++-- 4 files changed, 7 insertions(+), 24 deletions(-) diff --git a/swift/infer_engine/protocol.py b/swift/infer_engine/protocol.py index 4eb3873d40..8fd25ef483 100644 --- a/swift/infer_engine/protocol.py +++ b/swift/infer_engine/protocol.py @@ -156,7 +156,8 @@ class RolloutInferRequest(InferRequest): images: List[str] = field(default_factory=list) data_dict: Dict = field(default_factory=dict) uuid: Optional[str] = None - add_eos: Optional[bool] = 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 8d57dcb30c..44c25fca8c 100644 --- a/swift/ray/megatron/gkd_trainer.py +++ b/swift/ray/megatron/gkd_trainer.py @@ -11,7 +11,6 @@ from swift.infer_engine.protocol import RequestConfig, RolloutOutput from swift.rl_core.data import GKDSample -from swift.rlhf_trainers.gkd_helpers import set_teacher_request_eos from swift.rlhf_trainers.gkd_loss import DataSource, TeacherOutput from swift.rlhf_trainers.utils import parse_prompt_logprobs from swift.rollout import MultiTurnScheduler, invoke_async_hook, multi_turns, run_multi_turn @@ -326,7 +325,7 @@ def _fetch_teacher_from_replicas(self, gkd_samples: List[GKDSample], samples): teacher_encodeds.append(teacher_encoded) else: teacher_encodeds.append(None) - set_teacher_request_eos(req, s, self.template) + 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/rlhf_trainers/gkd_helpers.py b/swift/rlhf_trainers/gkd_helpers.py index b57db81826..c91c363e90 100644 --- a/swift/rlhf_trainers/gkd_helpers.py +++ b/swift/rlhf_trainers/gkd_helpers.py @@ -7,12 +7,11 @@ """ import copy import torch -from dataclasses import asdict, dataclass, field +from dataclasses import dataclass, field from typing import Any, Dict, List, Optional, Tuple from swift.rl_core.data import GKDSample, OnPolicySample from swift.template.base import Template -from swift.template.template_inputs import StdTemplateInputs from swift.utils import get_cu_seqlens_from_position_ids, get_logger, json_parse_to_dict from .gkd_loss import TeacherOutput from .utils import (assemble_teacher_topk_logprobs, encode_sample, get_response_prefix_ids, @@ -81,22 +80,6 @@ def encode_gkd_samples( return student_encoded_list, teacher_encoded_list, has_opsd -def set_teacher_request_eos(request, sample: GKDSample, template: Optional[Template] = None) -> None: - """Resolve the training suffix policy before sending a GKD teacher request.""" - request.add_eos = sample.add_eos - if request.add_eos is not None or template is None or template.template_backend != 'swift': - return - # Render on a copy to reuse model-specific stop/overlap rules without changing - # the shared inference template or running multimodal processors again. - template = copy.copy(template) - template.set_mode('train') - inputs = StdTemplateInputs.from_dict(asdict(request)) - template._swift_prepare_inputs(inputs) - _, _, answer_len = template._swift_encode(inputs) - # The final answer consists of its response plus any appended suffix contexts. - request.add_eos = answer_len > 1 - - def build_teacher_requests(samples: List[OnPolicySample], template: Optional[Template] = None) -> List[Any]: """Build teacher API requests from samples (GKD or GRPO/OPD-RL). @@ -135,7 +118,7 @@ def build_teacher_requests(samples: List[OnPolicySample], template: Optional[Tem non_thinking_prefix_ids=prefix_ids) req.messages = messages if isinstance(s, GKDSample): - set_teacher_request_eos(req, s, template) + 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