Update rl_replacements.py
This commit is contained in:
parent
31df063c28
commit
34c66f0327
1 changed files with 4 additions and 5 deletions
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue