Distillflow 0.2.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (39) hide show
  1. distillflow/common/__init__.py +10 -0
  2. distillflow/common/common.py +98 -0
  3. distillflow/common/logger.py +101 -0
  4. distillflow/config/__init__.py +6 -0
  5. distillflow/config/config.py +26 -0
  6. distillflow/config/validator.py +38 -0
  7. distillflow/datasets/__init__.py +6 -0
  8. distillflow/datasets/args.py +99 -0
  9. distillflow/datasets/loader.py +191 -0
  10. distillflow/datasets/template/__init__.py +12 -0
  11. distillflow/datasets/template/alpaca.py +85 -0
  12. distillflow/datasets/template/args.py +56 -0
  13. distillflow/datasets/template/role.py +10 -0
  14. distillflow/datasets/template/sharegpt.py +84 -0
  15. distillflow/datasets/template/template.py +6 -0
  16. distillflow/evaluation/__init__.py +0 -0
  17. distillflow/evaluation/rouge.py +30 -0
  18. distillflow/model/__init__.py +0 -0
  19. distillflow/model/adapter.py +257 -0
  20. distillflow/model/args.py +173 -0
  21. distillflow/model/checkpoint.py +146 -0
  22. distillflow/model/finetuning_args.py +132 -0
  23. distillflow/model/generating_args.py +72 -0
  24. distillflow/model/liger_kernel.py +46 -0
  25. distillflow/model/loader.py +286 -0
  26. distillflow/model/quantization.py +193 -0
  27. distillflow/model/tokenizer.py +55 -0
  28. distillflow/model/unsloth.py +92 -0
  29. distillflow/trainer/AdaptationLayer.py +82 -0
  30. distillflow/trainer/__init__.py +0 -0
  31. distillflow/trainer/args.py +56 -0
  32. distillflow/trainer/attention_distillation.py +113 -0
  33. distillflow/trainer/fine_tuning.py +51 -0
  34. distillflow/trainer/layers_distillation.py +106 -0
  35. distillflow/trainer/logits_distillation.py +96 -0
  36. distillflow-0.2.0.dist-info/LICENSE +201 -0
  37. distillflow-0.2.0.dist-info/METADATA +162 -0
  38. distillflow-0.2.0.dist-info/RECORD +39 -0
  39. distillflow-0.2.0.dist-info/WHEEL +4 -0
@@ -0,0 +1,146 @@
1
+ import inspect
2
+ from functools import wraps, partial
3
+ from types import MethodType
4
+ from typing import Callable, Union, Any, Optional, Dict, Tuple
5
+
6
+ import torch
7
+ from transformers import PreTrainedModel
8
+
9
+ from .args import ModelArgs
10
+ from ..common import get_logger
11
+
12
+ logger = get_logger(__name__)
13
+
14
+ LAYERNORM_NAMES = {"norm", "ln"}
15
+
16
+ def get_unsloth_gradient_checkpointing_func() -> Callable:
17
+ class UnslothGradientCheckpointing(torch.autograd.Function):
18
+ r"""
19
+ Saves VRAM by smartly offloading to RAM.
20
+ """
21
+
22
+ @staticmethod
23
+ @torch.cuda.amp.custom_fwd
24
+ def forward(
25
+ ctx: torch.autograd.Function,
26
+ forward_function: torch.Module,
27
+ hidden_states: torch.Tensor,
28
+ *args: Union[torch.Tensor, Any],
29
+ ) -> torch.Tensor:
30
+ saved_hidden_states = hidden_states.to("cpu", non_blocking=True)
31
+ with torch.no_grad():
32
+ output = forward_function(hidden_states, *args)
33
+
34
+ ctx.save_for_backward(saved_hidden_states)
35
+ ctx.forward_function = forward_function
36
+ ctx.args = args
37
+ return output
38
+
39
+ @staticmethod
40
+ @torch.cuda.amp.custom_bwd
41
+ def backward(ctx: "torch.autograd.Function", grad_output: "torch.Tensor") -> "torch.Tensor":
42
+ (hidden_states,) = ctx.saved_tensors
43
+ hidden_states = hidden_states.to("cuda", non_blocking=True).detach()
44
+ hidden_states.requires_grad_(True)
45
+ with torch.enable_grad():
46
+ (output,) = ctx.forward_function(hidden_states, *ctx.args)
47
+
48
+ torch.autograd.backward(output, grad_output)
49
+ return (None, hidden_states.grad) + (None,) * len(ctx.args)
50
+
51
+ return UnslothGradientCheckpointing.apply
52
+
53
+
54
+ def get_custom_gradient_checkpointing_func(gradient_checkpointing_func: Callable) -> Callable:
55
+ r"""
56
+ Only applies gradient checkpointing to trainable layers.
57
+ """
58
+
59
+ @wraps(gradient_checkpointing_func)
60
+ def custom_gradient_checkpointing_func(func: Callable, *args: Union["torch.Tensor", Any], **kwargs):
61
+ underlying = func.func if isinstance(func, partial) else func
62
+ module: "torch.nn.Module" = underlying.__self__
63
+
64
+ if any(param.requires_grad for param in module.parameters()):
65
+ for arg in args:
66
+ if torch.is_tensor(arg) and torch.is_floating_point(arg):
67
+ arg.requires_grad_(True)
68
+
69
+ return gradient_checkpointing_func(func, *args, **kwargs)
70
+
71
+ if hasattr(gradient_checkpointing_func, "__self__"): # fix unsloth gc test case
72
+ custom_gradient_checkpointing_func.__self__ = gradient_checkpointing_func.__self__
73
+
74
+ return custom_gradient_checkpointing_func
75
+
76
+
77
+ def _gradient_checkpointing_enable(
78
+ self: PreTrainedModel,
79
+ gradient_checkpointing_kwargs: Optional[Dict[str, Any]] = None,
80
+ use_unsloth_gc: bool = False,
81
+ ) -> None:
82
+ r"""
83
+ Activates gradient checkpointing for the current model.
84
+
85
+ Modification of the original method to enable gradient checkpointing for block-wise optimizer.
86
+ """
87
+ from torch.utils.checkpoint import checkpoint
88
+
89
+ if not self.supports_gradient_checkpointing:
90
+ raise ValueError("{} does not support gradient checkpointing.".format(self.__class__.__name__))
91
+
92
+ if gradient_checkpointing_kwargs is None:
93
+ gradient_checkpointing_kwargs = {"use_reentrant": True}
94
+
95
+ if use_unsloth_gc:
96
+ gradient_checkpointing_func = get_unsloth_gradient_checkpointing_func()
97
+ else:
98
+ gradient_checkpointing_func = partial(checkpoint, **gradient_checkpointing_kwargs)
99
+
100
+ gradient_checkpointing_func = get_custom_gradient_checkpointing_func(gradient_checkpointing_func)
101
+ if "value" in inspect.signature(self._set_gradient_checkpointing).parameters: # old GC format
102
+ self.apply(partial(self._set_gradient_checkpointing, value=True))
103
+ self.enable_input_require_grads()
104
+ logger.warning("You are using the old GC format, some features (e.g. BAdam) will be invalid.")
105
+ else: # have already enabled input require gradients
106
+ self._set_gradient_checkpointing(enable=True, gradient_checkpointing_func=gradient_checkpointing_func)
107
+
108
+
109
+ def _fp32_forward_post_hook(
110
+ module: torch.nn.Module, args: Tuple[torch.Tensor], output: torch.Tensor
111
+ ) -> torch.Tensor:
112
+ return output.to(torch.float32)
113
+
114
+
115
+ def prepare_model_for_training(model: "PreTrainedModel", model_args: ModelArgs) -> None:
116
+ r"""
117
+ Includes:
118
+ (1) cast the layernorm in fp32
119
+ (2) make output embedding layer require grads
120
+ (3) add the upcasting of the lm_head in fp32
121
+ """
122
+ if model_args.upcast_layernorm:
123
+ logger.info("Upcasting layernorm weights in float32.")
124
+ for name, param in model.named_parameters():
125
+ if param.ndim == 1 and any(ln_name in name for ln_name in LAYERNORM_NAMES):
126
+ param.data = param.data.to(torch.float32)
127
+
128
+ if not model_args.disable_gradient_checkpointing:
129
+ if not getattr(model, "supports_gradient_checkpointing", False):
130
+ logger.warning("Current model does not support gradient checkpointing.")
131
+ else:
132
+ # use_reentrant=False might increase VRAM usage (have not been empirically verified yet)
133
+ # According to: https://github.com/huggingface/transformers/issues/28339
134
+ gradient_checkpointing_enable = partial(
135
+ _gradient_checkpointing_enable, use_unsloth_gc=model_args.use_unsloth_gc
136
+ )
137
+ model.gradient_checkpointing_enable = MethodType(gradient_checkpointing_enable, model)
138
+ model.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": True})
139
+ setattr(model.config, "use_cache", False) # turn off when gradient checkpointing is enabled
140
+ logger.info("Gradient checkpointing enabled.")
141
+
142
+ if model_args.upcast_lmhead_output:
143
+ output_layer = model.get_output_embeddings()
144
+ if isinstance(output_layer, torch.nn.Linear) and output_layer.weight.dtype != torch.float32:
145
+ logger.info("Upcasting lm_head outputs in float32.")
146
+ output_layer.register_forward_hook(_fp32_forward_post_hook)
@@ -0,0 +1,132 @@
1
+ from typing import Literal, List, Optional
2
+
3
+ from pydantic import Field, BaseModel
4
+
5
+
6
+ class FreezeArgs(BaseModel):
7
+ r"""
8
+ Arguments pertaining to the freeze (partial-parameter) training.
9
+ """
10
+
11
+ freeze_trainable_layers: int = Field(
12
+ default=2,
13
+ description=(
14
+ "The number of trainable layers for freeze (partial-parameter) fine-tuning. "
15
+ "Positive numbers mean the last n layers are set as trainable, "
16
+ "negative numbers mean the first n layers are set as trainable."
17
+ )
18
+ )
19
+ freeze_trainable_modules: str = Field(
20
+ default="all",
21
+ description=(
22
+ "Name(s) of trainable modules for freeze (partial-parameter) fine-tuning. "
23
+ "Use commas to separate multiple modules. "
24
+ "Use `all` to specify all the available modules."
25
+ )
26
+ )
27
+ freeze_extra_modules: Optional[str] = Field(
28
+ default=None,
29
+ description=(
30
+ "Name(s) of modules apart from hidden layers to be set as trainable "
31
+ "for freeze (partial-parameter) fine-tuning. "
32
+ "Use commas to separate multiple modules."
33
+ )
34
+ )
35
+
36
+
37
+ class LoraArgs(BaseModel):
38
+ r"""
39
+ Arguments pertaining to the LoRA training.
40
+ """
41
+
42
+ additional_target: Optional[str] = Field(
43
+ default=None,
44
+ description=(
45
+ "Name(s) of modules apart from LoRA layers to be set as trainable "
46
+ "and saved in the final checkpoint. "
47
+ "Use commas to separate multiple modules."
48
+ )
49
+ )
50
+ lora_alpha: Optional[int] = Field(
51
+ default=None,
52
+ description="The scale factor for LoRA fine-tuning (default: lora_rank * 2)."
53
+ )
54
+ lora_dropout: float = Field(
55
+ default=0.0,
56
+ description="Dropout rate for the LoRA fine-tuning."
57
+ )
58
+ lora_rank: int = Field(
59
+ default=8,
60
+ description="The intrinsic dimension for LoRA fine-tuning."
61
+ )
62
+ lora_target: str = Field(
63
+ default="all",
64
+ description=(
65
+ "Name(s) of target modules to apply LoRA. "
66
+ "Use commas to separate multiple modules. "
67
+ "Use `all` to specify all the linear modules."
68
+ )
69
+ )
70
+ loraplus_lr_ratio: Optional[float] = Field(
71
+ default=None,
72
+ description="LoRA plus learning rate ratio (lr_B / lr_A)."
73
+ )
74
+ loraplus_lr_embedding: float = Field(
75
+ default=1e-6,
76
+ description="LoRA plus learning rate for lora embedding layers."
77
+ )
78
+ use_rslora: bool = Field(
79
+ default=False,
80
+ description="Whether or not to use the rank stabilization scaling factor for LoRA layer."
81
+ )
82
+ use_dora: bool = Field(
83
+ default=False,
84
+ description="Whether or not to use the weight-decomposed lora method (DoRA)."
85
+ )
86
+ pissa_init: bool = Field(
87
+ default=False,
88
+ description="Whether or not to initialize a PiSSA adapter."
89
+ )
90
+ pissa_iter: int = Field(
91
+ default=16,
92
+ description="The number of iteration steps performed by FSVD in PiSSA. Use -1 to disable it."
93
+ )
94
+ pissa_convert: bool = Field(
95
+ default=False,
96
+ description="Whether or not to convert the PiSSA adapter to a normal LoRA adapter."
97
+ )
98
+ create_new_adapter: bool = Field(
99
+ default=False,
100
+ description="Whether or not to create a new adapter with randomly initialized weight."
101
+ )
102
+
103
+
104
+ class FinetuningArgs(FreezeArgs, LoraArgs, BaseModel):
105
+ r"""
106
+ Arguments pertaining to which techniques we are going to fine-tuning with.
107
+ """
108
+
109
+ finetuning_type: Literal["lora", "freeze", "full"] = Field(
110
+ default="full",
111
+ description="Which fine-tuning method to use."
112
+ )
113
+ use_llama_pro: bool = Field(
114
+ default=False,
115
+ description="Whether or not to make only the parameters in the expanded blocks trainable."
116
+ )
117
+
118
+ def __post_init__(self):
119
+ def split_arg(arg):
120
+ if isinstance(arg, str):
121
+ return [item.strip() for item in arg.split(",")]
122
+ return arg
123
+
124
+ self.freeze_trainable_modules: List[str] = split_arg(self.freeze_trainable_modules)
125
+ self.freeze_extra_modules: Optional[List[str]] = split_arg(self.freeze_extra_modules)
126
+ self.lora_alpha: int = self.lora_alpha or self.lora_rank * 2
127
+ self.lora_target: List[str] = split_arg(self.lora_target)
128
+ self.additional_target: Optional[List[str]] = split_arg(self.additional_target)
129
+
130
+ assert self.finetuning_type in ["lora", "freeze", "full"], "Invalid fine-tuning method."
131
+ if self.use_llama_pro and self.finetuning_type == "full":
132
+ raise ValueError("`use_llama_pro` is only valid for Freeze or LoRA training.")
@@ -0,0 +1,72 @@
1
+ from dataclasses import asdict, dataclass, field
2
+ from typing import Any, Dict, Optional
3
+
4
+ from transformers import GenerationConfig
5
+
6
+ @dataclass
7
+ class GeneratingArguments:
8
+ r"""
9
+ Arguments pertaining to specify the decoding parameters.
10
+ """
11
+
12
+ do_sample: bool = field(
13
+ default=True,
14
+ metadata={"help": "Whether or not to use sampling, use greedy decoding otherwise."},
15
+ )
16
+ temperature: float = field(
17
+ default=0.95,
18
+ metadata={"help": "The value used to modulate the next token probabilities."},
19
+ )
20
+ top_p: float = field(
21
+ default=0.7,
22
+ metadata={
23
+ "help": "The smallest set of most probable tokens with probabilities that add up to top_p or higher are kept."
24
+ },
25
+ )
26
+ top_k: int = field(
27
+ default=50,
28
+ metadata={"help": "The number of highest probability vocabulary tokens to keep for top-k filtering."},
29
+ )
30
+ num_beams: int = field(
31
+ default=1,
32
+ metadata={"help": "Number of beams for beam search. 1 means no beam search."},
33
+ )
34
+ max_length: int = field(
35
+ default=1024,
36
+ metadata={"help": "The maximum length the generated tokens can have. It can be overridden by max_new_tokens."},
37
+ )
38
+ max_new_tokens: int = field(
39
+ default=1024,
40
+ metadata={"help": "The maximum numbers of tokens to generate, ignoring the number of tokens in the prompt."},
41
+ )
42
+ repetition_penalty: float = field(
43
+ default=1.0,
44
+ metadata={"help": "The parameter for repetition penalty. 1.0 means no penalty."},
45
+ )
46
+ length_penalty: float = field(
47
+ default=1.0,
48
+ metadata={"help": "Exponential penalty to the length that is used with beam-based generation."},
49
+ )
50
+ default_system: Optional[str] = field(
51
+ default=None,
52
+ metadata={"help": "Default system message to use in chat completion."},
53
+ )
54
+ skip_special_tokens: bool = field(
55
+ default=True,
56
+ metadata={"help": "Whether or not to remove special tokens in the decoding."},
57
+ )
58
+
59
+ def to_dict(self, obey_generation_config: bool = False) -> Dict[str, Any]:
60
+ args = asdict(self)
61
+ if args.get("max_new_tokens", -1) > 0:
62
+ args.pop("max_length", None)
63
+ else:
64
+ args.pop("max_new_tokens", None)
65
+
66
+ if obey_generation_config:
67
+ generation_config = GenerationConfig()
68
+ for key in list(args.keys()):
69
+ if not hasattr(generation_config, key):
70
+ args.pop(key)
71
+
72
+ return args
@@ -0,0 +1,46 @@
1
+ import inspect
2
+ from transformers import PretrainedConfig
3
+ from ..common import logger
4
+
5
+ from distillflow.model.args import ModelArgs
6
+
7
+ logger = logger.get_logger(__name__)
8
+
9
+ def apply_liger_kernel(
10
+ config: PretrainedConfig,
11
+ model_args: ModelArgs,
12
+ is_trainable: bool,
13
+ require_logits: bool,
14
+ ) -> None:
15
+ if not is_trainable or not model_args.enable_liger_kernel:
16
+ return
17
+
18
+ model_type = getattr(config, "model_type", None)
19
+ if model_type == "gemma":
20
+ from liger_kernel.transformers import apply_liger_kernel_to_gemma as apply_liger_kernel
21
+ elif model_type == "gemma2":
22
+ from liger_kernel.transformers import apply_liger_kernel_to_gemma2 as apply_liger_kernel
23
+ elif model_type == "llama":
24
+ from liger_kernel.transformers import apply_liger_kernel_to_llama as apply_liger_kernel
25
+ elif model_type == "mistral":
26
+ from liger_kernel.transformers import apply_liger_kernel_to_mistral as apply_liger_kernel
27
+ elif model_type == "mixtral":
28
+ from liger_kernel.transformers import apply_liger_kernel_to_mixtral as apply_liger_kernel
29
+ elif model_type == "phi3":
30
+ from liger_kernel.transformers import apply_liger_kernel_to_phi3 as apply_liger_kernel
31
+ elif model_type == "qwen2":
32
+ from liger_kernel.transformers import apply_liger_kernel_to_qwen2 as apply_liger_kernel
33
+ elif model_type == "qwen2_vl":
34
+ from liger_kernel.transformers import apply_liger_kernel_to_qwen2_vl as apply_liger_kernel
35
+ else:
36
+ logger.warning("Current model does not support liger kernel.")
37
+ return
38
+
39
+ if require_logits and "fused_linear_cross_entropy" in inspect.signature(apply_liger_kernel).parameters:
40
+ logger.info("Current training stage does not support chunked cross entropy.")
41
+ kwargs = {"fused_linear_cross_entropy": False}
42
+ else:
43
+ kwargs = {}
44
+
45
+ apply_liger_kernel(**kwargs)
46
+ logger.info("Liger kernel has been applied to the model.")
@@ -0,0 +1,286 @@
1
+ import math
2
+ import os
3
+ from contextlib import nullcontext
4
+ from types import MethodType
5
+ from typing import Dict, Any
6
+
7
+ from transformers import PreTrainedTokenizer, AutoConfig, PreTrainedModel, \
8
+ PretrainedConfig, AutoModelForCausalLM, is_torch_npu_available
9
+ from transformers.integrations import is_deepspeed_zero3_enabled, is_deepspeed_available
10
+ from transformers.modeling_utils import is_fsdp_enabled
11
+ from transformers.utils import is_torch_sdpa_available, is_flash_attn_2_available
12
+ from transformers.utils.versions import require_version
13
+
14
+ from .adapter import init_adapter
15
+ from .checkpoint import prepare_model_for_training
16
+ from .liger_kernel import apply_liger_kernel
17
+ from .quantization import _configure_quantization, QuantizationMethod
18
+ from .tokenizer import load_tokenizer
19
+ from .unsloth import load_unsloth_pretrained_model
20
+ from ..common.logger import get_logger
21
+ from .args import ModelArgs
22
+ import torch
23
+
24
+ from ..common import count_parameters, infer_optim_dtype, get_current_device
25
+
26
+ logger = get_logger(__name__)
27
+
28
+ def get_init_kwargs(model_args: ModelArgs) -> Dict[str, Any]:
29
+ r"""
30
+ Gets arguments to load config/tokenizer/model.
31
+
32
+ Note: including inplace operation of model_args.
33
+ """
34
+ return {
35
+ "trust_remote_code": True,
36
+ "cache_dir": model_args.cache_dir,
37
+ "revision": model_args.model_revision,
38
+ "token": model_args.hf_hub_token
39
+ }
40
+
41
+ def _register_autoclass(config: PretrainedConfig, model: "PreTrainedModel", tokenizer: "PreTrainedTokenizer"):
42
+ if "AutoConfig" in getattr(config, "auto_map", {}):
43
+ config.__class__.register_for_auto_class()
44
+ if "AutoModelForCausalLM" in getattr(config, "auto_map", {}):
45
+ model.__class__.register_for_auto_class()
46
+ if "AutoTokenizer" in tokenizer.init_kwargs.get("auto_map", {}):
47
+ tokenizer.__class__.register_for_auto_class()
48
+
49
+ def _configure_attn_implementation(
50
+ config: PretrainedConfig, model_args: ModelArgs, is_trainable: bool
51
+ ) -> None:
52
+ if getattr(config, "model_type", None) == "gemma2" and is_trainable:
53
+ if model_args.flash_attn == "auto" or model_args.flash_attn == "fa2":
54
+ if is_flash_attn_2_available():
55
+ require_version("flash_attn>=2.6.3", "To fix: pip install flash_attn>=2.6.3")
56
+ if model_args.flash_attn != "fa2":
57
+ logger.warning("Gemma-2 should use flash attention 2, change `flash_attn` to fa2.")
58
+ model_args.flash_attn = "fa2"
59
+ else:
60
+ logger.warning("FlashAttention-2 is not installed, use eager attention.")
61
+ model_args.flash_attn = "disabled"
62
+ elif model_args.flash_attn == "sdpa":
63
+ logger.warning("Gemma-2 should use soft-capping attention, while the SDPA attention does not support it.")
64
+
65
+ if model_args.flash_attn == "auto":
66
+ return
67
+
68
+ elif model_args.flash_attn == "disabled":
69
+ requested_attn_implementation = "eager"
70
+
71
+ elif model_args.flash_attn == "sdpa":
72
+ if not is_torch_sdpa_available():
73
+ logger.warning("torch>=2.1.1 is required for SDPA attention.")
74
+ return
75
+
76
+ requested_attn_implementation = "sdpa"
77
+ elif model_args.flash_attn == "fa2":
78
+ if not is_flash_attn_2_available():
79
+ logger.warning("FlashAttention-2 is not installed.")
80
+ return
81
+
82
+ requested_attn_implementation = "flash_attention_2"
83
+ else:
84
+ raise NotImplementedError("Unknown attention type: {}".format(model_args.flash_attn))
85
+
86
+ if getattr(config, "model_type", None) == "internlm2": # special case for custom models
87
+ setattr(config, "attn_implementation", requested_attn_implementation)
88
+ else:
89
+ setattr(config, "_attn_implementation", requested_attn_implementation)
90
+
91
+ def _patch_config(
92
+ config: PretrainedConfig,
93
+ tokenizer: PreTrainedTokenizer,
94
+ model_args: ModelArgs,
95
+ init_kwargs: Dict[str, Any],
96
+ is_trainable: bool,
97
+ torch_dtype: torch.dtype
98
+ ) -> None:
99
+ if is_torch_npu_available():
100
+ use_jit_compile = os.environ.get("JIT_COMPILE", "0").lower() in ["true", "1"]
101
+ torch.npu.set_compile_mode(jit_compile=use_jit_compile)
102
+
103
+ _configure_attn_implementation(config, model_args, is_trainable)
104
+ _configure_quantization(config, tokenizer, model_args.quantization_args, init_kwargs, torch_dtype)
105
+
106
+ if model_args.use_cache and not is_trainable:
107
+ setattr(config, "use_cache", True)
108
+ logger.info("Using KV cache for faster generation.")
109
+
110
+ if getattr(config, "model_type", None) == "qwen":
111
+ setattr(config, "use_flash_attn", model_args.flash_attn == "fa2")
112
+ for dtype_name, dtype in [("fp16", torch.float16), ("bf16", torch.bfloat16), ("fp32", torch.float32)]:
113
+ setattr(config, dtype_name, torch_dtype == dtype)
114
+
115
+ if getattr(config, "model_type", None) == "qwen2" and is_trainable and model_args.flash_attn == "fa2":
116
+ setattr(config, "use_cache", False) # qwen2 does not support use_cache when using flash attn
117
+
118
+ # if "LlavaLlamaForCausalLM" in getattr(config, "architectures", []):
119
+ # raise ValueError("Please download llava models with hf-compatible format: https://huggingface.co/llava-hf")
120
+
121
+ # deepspeed zero3 is not compatible with low_cpu_mem_usage
122
+
123
+ init_kwargs["low_cpu_mem_usage"] = model_args.low_cpu_mem_usage and (not is_deepspeed_available())
124
+
125
+ # cast data type of the model if:
126
+ # 1. not deepspeed zero3 and not fsdp (keep zero3 or fsdp in float32)
127
+ # 2. quantization_bit is not None (qlora)
128
+
129
+ quantization_args = model_args.quantization_args
130
+ if (not is_deepspeed_zero3_enabled() and not is_fsdp_enabled()) or quantization_args.quantization_bit is not None:
131
+ init_kwargs["torch_dtype"] = torch_dtype
132
+
133
+ if init_kwargs["low_cpu_mem_usage"]: # device map requires low_cpu_mem_usage=True
134
+ if "device_map" not in init_kwargs:
135
+ init_kwargs["device_map"] = {"": get_current_device()}
136
+
137
+ if init_kwargs.get("device_map", None) == "auto":
138
+ init_kwargs["offload_folder"] = model_args.offload_folder
139
+
140
+
141
+ def _noisy_mean_initialization(embed_weight: "torch.Tensor", num_new_tokens: int) -> None:
142
+ embedding_dim = embed_weight.size(1)
143
+ avg_weight = embed_weight[:-num_new_tokens].mean(dim=0, keepdim=True)
144
+ noise_weight = torch.empty_like(embed_weight[-num_new_tokens:])
145
+ noise_weight.normal_(mean=0, std=(1.0 / math.sqrt(embedding_dim)))
146
+ embed_weight[-num_new_tokens:] = avg_weight + noise_weight
147
+
148
+ def _resize_embedding_layer(model: "PreTrainedModel", tokenizer: "PreTrainedTokenizer") -> None:
149
+ r"""
150
+ Resize token embeddings.
151
+ """
152
+ # if is_deepspeed_zero3_enabled():
153
+ # import deepspeed # type: ignore
154
+ #
155
+ # params = [model.get_input_embeddings().weight]
156
+ # if model.get_output_embeddings() is not None and not model.config.tie_word_embeddings:
157
+ # params.append(model.get_output_embeddings().weight)
158
+ #
159
+ # context_maybe_zero3 = deepspeed.zero.GatheredParameters(params, modifier_rank=0)
160
+ # else:
161
+ context_maybe_zero3 = nullcontext()
162
+
163
+ with context_maybe_zero3:
164
+ current_embedding_size = model.get_input_embeddings().weight.size(0)
165
+
166
+ if len(tokenizer) > current_embedding_size:
167
+ if getattr(model, "quantization_method", None):
168
+ raise ValueError("Cannot resize embedding layers of a quantized model.")
169
+
170
+ if not isinstance(model.get_output_embeddings(), torch.nn.Linear):
171
+ raise ValueError("Current model does not support resizing embedding layers.")
172
+
173
+ model.resize_token_embeddings(len(tokenizer), pad_to_multiple_of=64)
174
+ with context_maybe_zero3:
175
+ new_embedding_size = model.get_input_embeddings().weight.size(0)
176
+ num_new_tokens = new_embedding_size - current_embedding_size
177
+ _noisy_mean_initialization(model.get_input_embeddings().weight.data, num_new_tokens)
178
+ _noisy_mean_initialization(model.get_output_embeddings().weight.data, num_new_tokens)
179
+
180
+ logger.info("Resized token embeddings from {} to {}.".format(current_embedding_size, new_embedding_size))
181
+
182
+ def _patch_model(
183
+ model: PreTrainedModel,
184
+ tokenizer: PreTrainedTokenizer,
185
+ model_args: ModelArgs,
186
+ is_trainable: bool,
187
+ ) -> None:
188
+ gen_config = model.generation_config # check and fix generation config
189
+ if not gen_config.do_sample and (
190
+ (gen_config.temperature is not None and gen_config.temperature != 1.0)
191
+ or (gen_config.top_p is not None and gen_config.top_p != 1.0)
192
+ or (gen_config.typical_p is not None and gen_config.typical_p != 1.0)
193
+ ):
194
+ gen_config.do_sample = True
195
+
196
+ if "GenerationMixin" not in str(model.generate.__func__):
197
+ model.generate = MethodType(PreTrainedModel.generate, model)
198
+
199
+ if model_args.resize_vocab:
200
+ _resize_embedding_layer(model, tokenizer)
201
+
202
+ if is_trainable:
203
+ prepare_model_for_training(model, model_args)
204
+ # autocast_projector_dtype(model, model_args)
205
+ # add_z3_leaf_module(model)
206
+
207
+ attn = getattr(model.config, "_attn_implementation", None)
208
+ if attn == "flash_attention_2":
209
+ logger.info("Using FlashAttention-2 for faster training and inference.")
210
+ elif attn == "sdpa":
211
+ logger.info("Using torch SDPA for faster training and inference.")
212
+
213
+ def load_model(
214
+ model_args: ModelArgs,
215
+ is_trainable: bool = False,
216
+ ) -> (PreTrainedModel, PreTrainedTokenizer):
217
+ tokenizer = load_tokenizer(model_args)
218
+ init_kwargs = get_init_kwargs(model_args)
219
+ config = AutoConfig.from_pretrained(model_args.model_name_or_path, **init_kwargs)
220
+ if model_args.infer_dtype != "auto" and not is_trainable:
221
+ torch_dtype = getattr(torch, model_args.infer_dtype)
222
+ else:
223
+ torch_dtype = infer_optim_dtype(model_dtype=getattr(config, "torch_dtype", None))
224
+
225
+ _patch_config(config, tokenizer, model_args, init_kwargs, is_trainable, torch_dtype)
226
+
227
+ # More details here: https://github.com/linkedin/Liger-Kernel
228
+ apply_liger_kernel(config, model_args, is_trainable, require_logits=True)
229
+
230
+ model = None
231
+ lazy_load = False
232
+ if model_args.use_unsloth:
233
+ if model_args.adapter_name_or_path is not None:
234
+ lazy_load = True
235
+ elif is_trainable:
236
+ model = load_unsloth_pretrained_model(config, model_args)
237
+
238
+ quantization_args = model_args.quantization_args
239
+ if quantization_args.quantization_method == QuantizationMethod.GPTQ.value:
240
+ from auto_gptq import AutoGPTQForCausalLM, BaseQuantizeConfig
241
+ quantize_config = BaseQuantizeConfig(
242
+ bits=quantization_args.quantization_bit,
243
+ group_size=128, # Group size (optional, can be None)
244
+ desc_act=False # Disable activation descriptor (optional)
245
+ )
246
+ model = AutoGPTQForCausalLM.from_pretrained(model_args.model_name_or_path, bits=quantization_args.quantization_bit, group_size=128, quantize_config=quantize_config)
247
+
248
+ if model is None and not lazy_load:
249
+ init_kwargs["config"] = config
250
+ init_kwargs["pretrained_model_name_or_path"] = model_args.model_name_or_path
251
+ model = AutoModelForCausalLM.from_pretrained(**init_kwargs)
252
+
253
+ if not lazy_load:
254
+ _patch_model(model, tokenizer, model_args, is_trainable)
255
+ _register_autoclass(config, model, tokenizer)
256
+
257
+ model = init_adapter(config, model, model_args, is_trainable)
258
+ if not is_trainable:
259
+ model.requires_grad_(False)
260
+ for param in model.parameters():
261
+ if param.data.dtype == torch.float32 and torch_dtype != torch.float32:
262
+ param.data = param.data.to(torch_dtype)
263
+
264
+ model.eval()
265
+ else:
266
+ model.train()
267
+
268
+ trainable_params, all_param = count_parameters(model)
269
+ if is_trainable:
270
+ param_stats = "trainable params: {:,} || all params: {:,} || trainable%: {:.4f}".format(
271
+ trainable_params, all_param, 100 * trainable_params / all_param
272
+ )
273
+ else:
274
+ param_stats = "all params: {:,}".format(all_param)
275
+
276
+ logger.info(param_stats)
277
+
278
+ if model_args.print_param_status:
279
+ for name, param in model.named_parameters():
280
+ print(
281
+ "name: {}, dtype: {}, device: {}, trainable: {}".format(
282
+ name, param.dtype, param.device, param.requires_grad
283
+ )
284
+ )
285
+
286
+ return model, tokenizer