Update rl_replacements.py

This commit is contained in:
Daniel Han 2025-07-17 06:25:21 -07:00
commit 34c66f0327

View file

@ -256,11 +256,10 @@ def grpo_trainer__generate_and_score_completions(function_name, function):
replace_part, spacing = found[0]
removed_comments = re.sub(r"\#[^\n]{1,}", "", replace_part)
splits = removed_comments.split("\n")
if sum(re.search(rf"^{spacing}[^\s]", x) is not None for x in splits) == 2 and \
len(spacing) == 8:
replace_part
if sum(re.search(rf"^{spacing}[^\s]", x) is not None for x in splits) == 2 and len(spacing) == 8:
new_replacement = spacing + """if self.max_prompt_length is not None:
new_replacement = spacing + \
"""if self.max_prompt_length is not None:
# If max_prompt_length is set, we trim the prompt to keep only the last `max_prompt_length` tokens.
# Then we decode those tokens back into text. We manually remove leading pad tokens from the decoded text,
# because we can't use `skip_special_tokens=True` (some special tokens are still needed for generation).
@ -282,7 +281,7 @@ def grpo_trainer__generate_and_score_completions(function_name, function):
# Generate completions using either vLLM or regular generation
if self.use_vllm:"""
function = function.replace(replace_part, new_replacement)
function = function.replace(replace_part, new_replacement)
pass
return function
pass