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.
- distillflow/common/__init__.py +10 -0
- distillflow/common/common.py +98 -0
- distillflow/common/logger.py +101 -0
- distillflow/config/__init__.py +6 -0
- distillflow/config/config.py +26 -0
- distillflow/config/validator.py +38 -0
- distillflow/datasets/__init__.py +6 -0
- distillflow/datasets/args.py +99 -0
- distillflow/datasets/loader.py +191 -0
- distillflow/datasets/template/__init__.py +12 -0
- distillflow/datasets/template/alpaca.py +85 -0
- distillflow/datasets/template/args.py +56 -0
- distillflow/datasets/template/role.py +10 -0
- distillflow/datasets/template/sharegpt.py +84 -0
- distillflow/datasets/template/template.py +6 -0
- distillflow/evaluation/__init__.py +0 -0
- distillflow/evaluation/rouge.py +30 -0
- distillflow/model/__init__.py +0 -0
- distillflow/model/adapter.py +257 -0
- distillflow/model/args.py +173 -0
- distillflow/model/checkpoint.py +146 -0
- distillflow/model/finetuning_args.py +132 -0
- distillflow/model/generating_args.py +72 -0
- distillflow/model/liger_kernel.py +46 -0
- distillflow/model/loader.py +286 -0
- distillflow/model/quantization.py +193 -0
- distillflow/model/tokenizer.py +55 -0
- distillflow/model/unsloth.py +92 -0
- distillflow/trainer/AdaptationLayer.py +82 -0
- distillflow/trainer/__init__.py +0 -0
- distillflow/trainer/args.py +56 -0
- distillflow/trainer/attention_distillation.py +113 -0
- distillflow/trainer/fine_tuning.py +51 -0
- distillflow/trainer/layers_distillation.py +106 -0
- distillflow/trainer/logits_distillation.py +96 -0
- distillflow-0.2.0.dist-info/LICENSE +201 -0
- distillflow-0.2.0.dist-info/METADATA +162 -0
- distillflow-0.2.0.dist-info/RECORD +39 -0
- 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
|