From 99e07e7c905d937d94dc69db30e507995968072b Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Wed, 20 Nov 2024 03:36:41 -0800 Subject: [PATCH] patch_fast_lora --- unsloth/kernels/__init__.py | 1 + unsloth/kernels/fast_lora.py | 75 ++++++++++++++++++++++++++++++++++++ unsloth/models/_utils.py | 7 ++++ 3 files changed, 83 insertions(+) diff --git a/unsloth/kernels/__init__.py b/unsloth/kernels/__init__.py index 82e7641693..ef5fa5da70 100644 --- a/unsloth/kernels/__init__.py +++ b/unsloth/kernels/__init__.py @@ -42,6 +42,7 @@ from .fast_lora import ( apply_lora_mlp_geglu_approx, apply_lora_qkv, apply_lora_o, + fast_lora_forward, ) from .utils import fast_dequantize, fast_gemv, QUANT_STATE, fast_linear_forward, matmul_lora diff --git a/unsloth/kernels/fast_lora.py b/unsloth/kernels/fast_lora.py index 2177b43b9e..6481661d88 100644 --- a/unsloth/kernels/fast_lora.py +++ b/unsloth/kernels/fast_lora.py @@ -410,3 +410,78 @@ def apply_lora_o(self, X): O = LoRA_W.apply(X, OW, OW_quant, OA, OB, OS) return O pass + + +IDENTITY_DROPOUT = torch.nn.Identity +@torch._disable_dynamo +def fast_lora_forward(self, x: torch.Tensor, *args, **kwargs) -> torch.Tensor: + self._check_forward_args(x, *args, **kwargs) + adapter_names = kwargs.pop("adapter_names", None) + + if self.disable_adapters: + if self.merged: + self.unmerge() + result = self.base_layer(x, *args, **kwargs) + elif adapter_names is not None: + result = self._mixed_batch_forward(x, *args, adapter_names=adapter_names, **kwargs) + elif self.merged: + result = self.base_layer(x, *args, **kwargs) + else: + # Fastpath + if len(self.active_adapters) == 1: + active_adapter = self.active_adapters[0] + if active_adapter not in self.lora_A.keys(): return self.base_layer(x, *args, **kwargs) + + dropout = self.lora_dropout[active_adapter] + if isinstance(dropout, IDENTITY_DROPOUT) and not self.use_dora[active_adapter]: + lora_A = self.lora_A[active_adapter].weight + lora_B = self.lora_B[active_adapter].weight + scaling = self.scaling[active_adapter] + W = self.base_layer.weight + return LoRA_W.apply(x, W, QUANT_STATE(W), lora_A, lora_B, scaling) + pass + pass + + result = self.base_layer(x, *args, **kwargs) + # As per Tim Dettmers, for 4bit, we need to defensively clone here. + # The reason is that in some cases, an error can occur that backprop + # does not work on a manipulated view. This issue may be solved with + # newer PyTorch versions but this would need extensive testing to be + # sure. + result = result.clone() + + for active_adapter in self.active_adapters: + if active_adapter not in self.lora_A.keys(): + continue + lora_A = self.lora_A[active_adapter] + lora_B = self.lora_B[active_adapter] + dropout = self.lora_dropout[active_adapter] + scaling = self.scaling[active_adapter] + + requires_conversion = not torch.is_autocast_enabled() + if requires_conversion: + expected_dtype = result.dtype + x = x.to(lora_A.weight.dtype) + + if not self.use_dora[active_adapter]: + result = result + lora_B(lora_A(dropout(x))) * scaling + else: + if isinstance(dropout, torch.nn.Identity) or not self.training: + base_result = result + else: + x = dropout(x) + base_result = None + + result = result + self.lora_magnitude_vector[active_adapter]( + x, + lora_A=lora_A, + lora_B=lora_B, + scaling=scaling, + base_layer=self.get_base_layer(), + base_result=base_result, + ) + if requires_conversion: + result = result.to(expected_dtype) + + return result +pass diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index a93c443042..4271eb6a82 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -62,6 +62,7 @@ __all__ = [ "patch_compiled_autograd", "process_vision_info", "unsloth_compile_transformers", + "patch_fast_lora", ] import torch @@ -1081,6 +1082,12 @@ def patch_tokenizer(model, tokenizer): pass +def patch_fast_lora(): + import peft.tuners.lora.bnb + peft.tuners.lora.bnb.Linear4bit.forward = fast_lora_forward +pass + + def unsloth_compile_transformers( model_name, token = None,