Update rl.py
This commit is contained in:
parent
8506d6be95
commit
aebeeb4901
1 changed files with 7 additions and 9 deletions
|
|
@ -39,15 +39,13 @@ def PatchRL(FastLanguageModel):
|
|||
@contextmanager
|
||||
def unsloth_unwrap_model_for_generation(model, accelerator):
|
||||
# Must use for_inference to allow inference in Unsloth
|
||||
with torch.inference_mode():
|
||||
with unwrap_model_for_generation(model, accelerator) as unwrapped_model:
|
||||
FastLanguageModel.for_inference(unwrapped_model)
|
||||
try:
|
||||
yield unwrapped_model
|
||||
finally:
|
||||
# Finally return back training
|
||||
FastLanguageModel.for_training(model)
|
||||
pass
|
||||
with unwrap_model_for_generation(model, accelerator) as unwrapped_model:
|
||||
FastLanguageModel.for_inference(unwrapped_model)
|
||||
try:
|
||||
yield unwrapped_model.eval()
|
||||
finally:
|
||||
# Finally return back training
|
||||
FastLanguageModel.for_training(model)
|
||||
pass
|
||||
pass
|
||||
pass
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue