Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions swift/infer_engine/protocol.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
1 change: 1 addition & 0 deletions swift/ray/megatron/gkd_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
5 changes: 3 additions & 2 deletions swift/rl_core/data.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
Expand Down Expand Up @@ -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
Expand Down
2 changes: 2 additions & 0 deletions swift/rlhf_trainers/gkd_helpers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
4 changes: 2 additions & 2 deletions swift/template/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading