From e269b0dc59afee09c8c8d9544fd4c032f8690526 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Tue, 19 Mar 2024 19:55:02 +1100 Subject: [PATCH] Fix Saving (#264) * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * llama * Update llama.py * gemma * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update save.py * RoPE * Update llama.py * Update llama.py * Update llama.py * Update gemma.py * correct_dtype * Update gemma.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Chat Templates * Update README.md * Update README.md * Update llama.py * DoRA * Update _utils.py * Update chat_templates.py * Update llama.py * Hotfix - fix DoRA, Gemma prompt template (#202) (#203) * Update save.py * saving * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update __init__.py * Update save.py * Update save.py * Update save.py * save * trainer * spaces * original * Gemma * Update pyproject.toml * Update mapper.py * Update fast_lora.py * FastGemmaModel * model_type * Update llama.py * Update llama.py * Update gemma.py * Update gemma.py * Update gemma.py * Update llama.py * Update llama.py * Update fast_lora.py * Update llama.py * Update llama.py * Update cross_entropy_loss.py * Update llama.py * Update llama.py * gemma * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update fast_lora.py * Update fast_lora.py * Fast CE Loss * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * CE * Update llama.py * Update llama.py * Update cross_entropy_loss.py * Update geglu.py * Update cross_entropy_loss.py * revert * Update llama.py * Update llama.py * norm * Update gemma.py * Update gemma.py * position_ids * Update gemma.py * Update gemma.py * pos * Update llama.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update cross_entropy_loss.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update llama.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update llama.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * revert * revert * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update llama.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update cross_entropy_loss.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * rope * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * llama * Update llama.py * gemma * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update save.py * RoPE * Update llama.py * Update llama.py * Update llama.py * Update gemma.py * correct_dtype * Update gemma.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Chat Templates * Update README.md * Update README.md * Update llama.py * DoRA * Update _utils.py * Update chat_templates.py * Update pyproject.toml * Small fixes * Update pyproject.toml * Approx gelu * Update geglu.py * Approx gelu * Update llama.py * Update __init__.py * Update __init__.py * Update _utils.py * Update geglu.py * Update gemma.py * Update rms_layernorm.py * Update rms_layernorm.py * Update rms_layernorm.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Fix Gemma merging * Update rms_layernorm.py * Update gemma.py * Update pyproject.toml * Layernorms * Gemma precision * Update gemma.py * sqrt * Update gemma.py * Update save.py * RoPE and Gemma precision * Update rms_layernorm.py * Fix warning * Update chat_templates.py * Update chat_templates.py * Update save.py * Update save.py * Update save.py * Update chat_templates.py * Update llama.py * model_name * Update loader.py * Tokenizer overwritten * Update llama.py * Update llama.py * Update llama.py * Update save.py * Accuracy * Revert * Update save.py * Update fast_lora.py * Update fast_lora.py * Update fast_lora.py * Update fast_lora.py * Update fast_lora.py * Update chat_templates.py * Update save.py * Update save.py * Update llama.py * Update llama.py * Account for DoRA * Update llama.py * Update save.py * GGUF incorrect * Update save.py * Update pyproject.toml * kaggle new * Update pyproject.toml * Update pyproject.toml * upcasting * Fix Colab * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update chat_templates.py * Update chat_templates.py * Update chat_templates.py * Update chat_templates.py * Update chat_templates.py * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update rope_embedding.py * Update rope_embedding.py * Fix bugs * Update fast_lora.py * Update fast_lora.py * Update README.md * Update README.md * GGUF * Update save.py * Update save.py * Update save.py * Update save.py * Update README.md * Update README.md * Bugs * Update fast_lora.py * Update pyproject.toml * Update fast_lora.py * Update __init__.py * Update fast_lora.py * dtype * Update llama.py * Update llama.py * Update llama.py * dtype * Update mistral.py * trust_remote_code * lm_head * Update llama.py * save_pretrained_settings * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * state_dict * Update save.py * whoami * Update llama.py * Update save.py --- unsloth/models/_utils.py | 86 ++++++++++++++++++++++++++++++++++++++++ unsloth/models/llama.py | 7 ++++ unsloth/save.py | 30 +++++++------- 3 files changed, 108 insertions(+), 15 deletions(-) diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index 02eea4ba0c..6f7da0f32a 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -26,6 +26,7 @@ from transformers import AutoTokenizer from platform import system as platform_system platform_system = platform_system() import math +import numpy as np __version__ = "2024.3" @@ -269,3 +270,88 @@ except: "Luckily, your training run will still work in the meantime!" ) pass + + +def _calculate_n_gradient_checkpoints( + n_layers : int, + method : Optional[Union[str, int]] = "sqrt", +) -> List[int]: + assert(type(n_layers) is int and n_layers > 0) + + if method is None: method = "sqrt" + + if method == "sqrt": + n_checkpoints = int(n_layers**0.5) + elif type(method) is int and method > 0: + n_checkpoints = int(np.ceil(n_layers / method)) + else: + raise ValueError("method must be 'sqrt' or an int >0 and <= n_layers.") + + size = n_layers // n_checkpoints + sizes = np.full(n_checkpoints, size, dtype = int) + leftovers = n_layers % n_checkpoints + # We append leftovers from the right + for k in range(leftovers): + sizes[n_checkpoints-1-k] += 1 + boundaries = np.hstack((0, np.cumsum(sizes))) + boundaries = boundaries.tolist() + return boundaries +pass + + +def calculate_n_gradient_checkpoints( + n_layers : int, + layers_per_checkpoint : Optional[Union[str, int]] = "sqrt", +) -> List[int]: + assert(type(n_layers) is int and n_layers > 0) + + if layers_per_checkpoint is None or layers_per_checkpoint == 1: + return None + + boundaries = _calculate_n_gradient_checkpoints(n_layers, layers_per_checkpoint) + + assert(boundaries[0] == 0 and boundaries[-1] == n_layers) + assert(min(boundaries) == 0 and max(boundaries) == n_layers) + assert(np.diff(boundaries).min() >= 0) + return boundaries +pass + + +def prepare_n_gradient_checkpoints( + model : Any, + layers_per_checkpoint : Optional[Union[str, int]] = "sqrt", + use_reentrant : Optional[bool] = True, +) -> None: + """ + Calculates where to place the gradient checkpoints given n_layers. + + Args: + model: Any LlamaModel with layers. + layers_per_checkpoint (`Union[str, int]`, *optional*): + Can either be `sqrt` or an integer for how many layers per checkpoint you want. + The more, the less memory usage, but can be slower. Default is `sqrt`. + Choose 1 for Pytorch gradient checkpointing. 2 to wrap 2 layers in 1 module etc. + use_reentrant (`bool`, *optional*): + https://github.com/pytorch/pytorch/blob/main/torch/utils/checkpoint.py#L354 + Optimal gradient checkpointing algorithm `use_reentrant=False` which will + be the default in future Pytorch versions doesn't seem to work?? + """ + _model = None + if hasattr(model, "layers"): + _model = model + elif hasattr(model, "model"): + if hasattr(model.model, "layers"): + _model = model.model + if _model is None: + raise TypeError("`model` or `model.model` does not have attribute `layers`. Are you sure this is a model?") + pass + + if use_reentrant is False: + use_reentrant = True + pass + + n_layers = len(_model.layers) + boundaries = calculate_n_gradient_checkpoints(n_layers, layers_per_checkpoint) + _model._gradient_checkpointing_boundaries = boundaries + _model._gradient_checkpointing_use_reentrant = use_reentrant +pass diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 574385fba0..b7aae06ee5 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -593,6 +593,13 @@ def LlamaModel_fast_forward( all_self_attns = () if output_attentions else None next_decoder_cache = () if use_cache else None + # Gradient checkpointing methods (ie sqrt) + if hasattr(self, "_gradient_checkpointing_boundaries"): + boundaries = self._gradient_checkpointing_boundaries + else: + boundaries = None + pass + for idx, decoder_layer in enumerate(self.layers): if output_hidden_states: all_hidden_states += (hidden_states,) diff --git a/unsloth/save.py b/unsloth/save.py index c8b4053a10..a161394057 100644 --- a/unsloth/save.py +++ b/unsloth/save.py @@ -24,6 +24,7 @@ from transformers.models.llama.modeling_llama import logger from .kernels import fast_dequantize, QUANT_STATE, get_lora_parameters import subprocess import psutil +import re __all__ = [ "print_quantization_methods", @@ -176,19 +177,6 @@ def unsloth_save_model( temporary_location : str = "_unsloth_temporary_saved_buffers", maximum_memory_usage : float = 0.9, ): - # First check for a token! - if push_to_hub: - from huggingface_hub import whoami - try: - username = whoami(token = token)["name"] - except: - raise RuntimeError( - "Unsloth: Please supply a token!\n"\ - "Go to https://huggingface.co/settings/tokens" - ) - pass - pass - if commit_message is None: commit_message = "" if "Unsloth" not in commit_message: commit_message += " (Trained with Unsloth)" @@ -215,7 +203,19 @@ def unsloth_save_model( for deletion in ("model", "tokenizer", "save_method", "temporary_location", "maximum_memory_usage"): del save_pretrained_settings[deletion] pass - import re + + # First check for a token! + if push_to_hub: + from huggingface_hub import whoami + try: + username = whoami(token = token)["name"] + except: + raise RuntimeError( + "Unsloth: Please supply a token!\n"\ + "Go to https://huggingface.co/settings/tokens" + ) + pass + pass assert(maximum_memory_usage > 0 and maximum_memory_usage <= 0.95) @@ -588,7 +588,7 @@ def unsloth_save_model( from huggingface_hub import HfApi hf_api = HfApi(token = save_pretrained_settings["token"]) - print("Unsloth: Uploading all files... Please wait!") + print("Unsloth: Uploading all files... Please wait...") hf_api.upload_folder( folder_path = new_save_directory, path_in_repo = ".",