Update _utils.py
This commit is contained in:
parent
d426961266
commit
51257d37a8
1 changed files with 8 additions and 1 deletions
|
|
@ -12,7 +12,7 @@
|
|||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
__version__ = "2024.12.12"
|
||||
__version__ = "2025.1.1"
|
||||
|
||||
__all__ = [
|
||||
"prepare_model_for_kbit_training",
|
||||
|
|
@ -110,6 +110,9 @@ from unsloth_zoo.compiler import (
|
|||
get_transformers_model_type,
|
||||
unsloth_compile_transformers as _unsloth_compile_transformers,
|
||||
)
|
||||
from unsloth_zoo.peft_utils import (
|
||||
requires_grad_for_gradient_checkpointing,
|
||||
)
|
||||
|
||||
# =============================================
|
||||
# Disable some warnings which can get annoying
|
||||
|
|
@ -557,6 +560,10 @@ def prepare_model_for_kbit_training(
|
|||
output.requires_grad_(True)
|
||||
model.get_input_embeddings().register_forward_hook(make_inputs_require_grad)
|
||||
|
||||
# Enable grads on non language models as well
|
||||
requires_grad_for_gradient_checkpointing()
|
||||
pass
|
||||
|
||||
return model
|
||||
pass
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue