Update rl.py
This commit is contained in:
parent
791c6fc04f
commit
8bab4a5eec
1 changed files with 8 additions and 8 deletions
|
|
@ -71,7 +71,14 @@ def PatchRL(FastLanguageModel):
|
|||
pass
|
||||
|
||||
|
||||
# Handles NEFTune
|
||||
RLTrainer_replacement = '''
|
||||
import os
|
||||
from typing import *
|
||||
from dataclasses import dataclass, field
|
||||
from packaging.version import Version
|
||||
import torch
|
||||
|
||||
# https://github.com/huggingface/transformers/blob/main/src/transformers/trainer_utils.py#L126
|
||||
def neftune_post_forward_hook(module, input, output):
|
||||
if module.training:
|
||||
dims = torch.tensor(output.size(1) * output.size(2))
|
||||
|
|
@ -80,13 +87,6 @@ def neftune_post_forward_hook(module, input, output):
|
|||
return output
|
||||
pass
|
||||
|
||||
|
||||
RLTrainer_replacement = '''
|
||||
import os
|
||||
from typing import *
|
||||
from dataclasses import dataclass, field
|
||||
from packaging.version import Version
|
||||
|
||||
@dataclass
|
||||
class Unsloth{RLConfig_name}({RLConfig_name}):
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue