Update _utils.py
This commit is contained in:
parent
7f9ecd592f
commit
020cbd2dd5
1 changed files with 14 additions and 1 deletions
|
|
@ -41,6 +41,8 @@ __all__ = [
|
|||
"torch_amp_custom_bwd",
|
||||
"accelerate_old_send_to_device",
|
||||
"accelerate_new_send_to_device",
|
||||
"patch_gradient_checkpointing",
|
||||
"unpatch_gradient_checkpointing",
|
||||
]
|
||||
|
||||
import torch
|
||||
|
|
@ -791,7 +793,7 @@ class Unsloth_Offloaded_Gradient_Checkpointer(torch.autograd.Function):
|
|||
def backward(ctx, dY):
|
||||
(hidden_states,) = ctx.saved_tensors
|
||||
hidden_states = hidden_states.to("cuda:0", non_blocking = True).detach()
|
||||
hidden_states.requires_grad = True
|
||||
hidden_states.requires_grad_(True)
|
||||
with torch.enable_grad():
|
||||
(output,) = ctx.forward_function(hidden_states, *ctx.args)
|
||||
torch.autograd.backward(output, dY)
|
||||
|
|
@ -806,6 +808,17 @@ def unsloth_offloaded_gradient_checkpoint(function, *args, use_reentrant = None,
|
|||
pass
|
||||
|
||||
|
||||
import torch.utils
|
||||
old_checkpoint = torch.utils.checkpoint
|
||||
def patch_gradient_checkpointing():
|
||||
torch.utils.checkpoint = unsloth_offloaded_gradient_checkpoint
|
||||
pass
|
||||
|
||||
def unpatch_gradient_checkpointing():
|
||||
torch.utils.checkpoint = old_checkpoint
|
||||
pass
|
||||
|
||||
|
||||
# =============================================
|
||||
# Fixes Bitsandbytes to remove missing warnings
|
||||
from transformers.utils.quantization_config import BitsAndBytesConfig, QuantizationMethod
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue