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,193 @@
|
|
|
1
|
+
import os
|
|
2
|
+
from enum import Enum, unique
|
|
3
|
+
from random import random
|
|
4
|
+
from typing import Any, Dict, Optional, List
|
|
5
|
+
|
|
6
|
+
import torch
|
|
7
|
+
from datasets import load_dataset
|
|
8
|
+
from transformers import BitsAndBytesConfig, EetqConfig, HqqConfig, GPTQConfig
|
|
9
|
+
from transformers.integrations import is_deepspeed_zero3_enabled
|
|
10
|
+
from transformers.modeling_utils import is_fsdp_enabled
|
|
11
|
+
from transformers.utils.versions import require_version
|
|
12
|
+
|
|
13
|
+
from .args import QuantizationArgs
|
|
14
|
+
from ..common.logger import get_logger
|
|
15
|
+
from ..common import get_current_device
|
|
16
|
+
|
|
17
|
+
from transformers import PretrainedConfig, PreTrainedTokenizer
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
logger = get_logger(__name__)
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
@unique
|
|
24
|
+
class QuantizationMethod(str, Enum):
|
|
25
|
+
r"""
|
|
26
|
+
Borrowed from `transformers.utils.quantization_config.QuantizationMethod`.
|
|
27
|
+
"""
|
|
28
|
+
|
|
29
|
+
BITS_AND_BYTES = "bitsandbytes"
|
|
30
|
+
GPTQ = "gptq"
|
|
31
|
+
AWQ = "awq"
|
|
32
|
+
AQLM = "aqlm"
|
|
33
|
+
QUANTO = "quanto"
|
|
34
|
+
EETQ = "eetq"
|
|
35
|
+
HQQ = "hqq"
|
|
36
|
+
|
|
37
|
+
FILEEXT2TYPE = {
|
|
38
|
+
"arrow": "arrow",
|
|
39
|
+
"csv": "csv",
|
|
40
|
+
"json": "json",
|
|
41
|
+
"jsonl": "json",
|
|
42
|
+
"parquet": "parquet",
|
|
43
|
+
"txt": "text",
|
|
44
|
+
}
|
|
45
|
+
|
|
46
|
+
def _get_quantization_dataset(tokenizer: "PreTrainedTokenizer", model_args: "ModelArguments") -> List[Dict[str, Any]]:
|
|
47
|
+
r"""
|
|
48
|
+
Prepares the tokenized dataset to perform AutoGPTQ. Do not use tensor output for JSON serialization.
|
|
49
|
+
"""
|
|
50
|
+
if os.path.isfile(model_args.export_quantization_dataset):
|
|
51
|
+
data_path = FILEEXT2TYPE.get(model_args.export_quantization_dataset.split(".")[-1], None)
|
|
52
|
+
data_files = model_args.export_quantization_dataset
|
|
53
|
+
else:
|
|
54
|
+
data_path = model_args.export_quantization_dataset
|
|
55
|
+
data_files = None
|
|
56
|
+
|
|
57
|
+
dataset = load_dataset(
|
|
58
|
+
path=data_path,
|
|
59
|
+
data_files=data_files,
|
|
60
|
+
split="train",
|
|
61
|
+
cache_dir=model_args.cache_dir,
|
|
62
|
+
token=model_args.hf_hub_token,
|
|
63
|
+
)
|
|
64
|
+
|
|
65
|
+
samples = []
|
|
66
|
+
maxlen = model_args.export_quantization_maxlen
|
|
67
|
+
for _ in range(model_args.export_quantization_nsamples):
|
|
68
|
+
n_try = 0
|
|
69
|
+
while True:
|
|
70
|
+
if n_try > 100:
|
|
71
|
+
raise ValueError("Cannot find satisfying example, considering decrease `export_quantization_maxlen`.")
|
|
72
|
+
|
|
73
|
+
sample_idx = random.randint(0, len(dataset) - 1)
|
|
74
|
+
sample: Dict[str, "torch.Tensor"] = tokenizer(dataset[sample_idx]["text"], return_tensors="pt")
|
|
75
|
+
n_try += 1
|
|
76
|
+
if sample["input_ids"].size(1) > maxlen:
|
|
77
|
+
break # TODO: fix large maxlen
|
|
78
|
+
|
|
79
|
+
word_idx = random.randint(0, sample["input_ids"].size(1) - maxlen - 1)
|
|
80
|
+
input_ids = sample["input_ids"][:, word_idx : word_idx + maxlen]
|
|
81
|
+
attention_mask = sample["attention_mask"][:, word_idx : word_idx + maxlen]
|
|
82
|
+
samples.append({"input_ids": input_ids.tolist(), "attention_mask": attention_mask.tolist()})
|
|
83
|
+
|
|
84
|
+
return samples
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
def _configure_quantization(
|
|
88
|
+
config: PretrainedConfig,
|
|
89
|
+
tokenizer: PreTrainedTokenizer,
|
|
90
|
+
quantization_args: QuantizationArgs,
|
|
91
|
+
init_kwargs: Dict[str, Any],
|
|
92
|
+
torch_dtype: Optional[torch.dtype]
|
|
93
|
+
) -> None:
|
|
94
|
+
r"""
|
|
95
|
+
Priority: PTQ-quantized (train/infer) > AutoGPTQ (export) > On-the-fly quantization (train/infer)
|
|
96
|
+
"""
|
|
97
|
+
|
|
98
|
+
if quantization_args is None:
|
|
99
|
+
return
|
|
100
|
+
if getattr(config, "quantization_config", None): # ptq
|
|
101
|
+
if quantization_args.quantization_bit is not None:
|
|
102
|
+
logger.warning("`quantization_bit` will not affect on the PTQ-quantized models.")
|
|
103
|
+
|
|
104
|
+
if is_deepspeed_zero3_enabled() or is_fsdp_enabled():
|
|
105
|
+
raise ValueError("DeepSpeed ZeRO-3 or FSDP is incompatible with PTQ-quantized models.")
|
|
106
|
+
|
|
107
|
+
quantization_config: Dict[str, Any] = getattr(config, "quantization_config", None)
|
|
108
|
+
quant_method = quantization_config.get("quant_method", "")
|
|
109
|
+
|
|
110
|
+
if quant_method == QuantizationMethod.GPTQ:
|
|
111
|
+
require_version("auto_gptq>=0.5.0", "To fix: pip install auto_gptq>=0.5.0")
|
|
112
|
+
quantization_config.pop("disable_exllama", None) # remove deprecated args
|
|
113
|
+
quantization_config["use_exllama"] = False # disable exllama
|
|
114
|
+
|
|
115
|
+
if quant_method == QuantizationMethod.AWQ:
|
|
116
|
+
require_version("autoawq", "To fix: pip install autoawq")
|
|
117
|
+
|
|
118
|
+
if quant_method == QuantizationMethod.AQLM:
|
|
119
|
+
require_version("aqlm>=1.1.0", "To fix: pip install aqlm[gpu]>=1.1.0")
|
|
120
|
+
quantization_config["bits"] = 2
|
|
121
|
+
|
|
122
|
+
quant_bits = quantization_config.get("bits", "?")
|
|
123
|
+
logger.info("Loading {}-bit {}-quantized model.".format(quant_bits, quant_method.upper()))
|
|
124
|
+
|
|
125
|
+
elif quantization_args.export_quantization_bit is not None: # auto-gptq
|
|
126
|
+
if quantization_args.export_quantization_bit not in [8, 4, 3, 2]:
|
|
127
|
+
raise ValueError("AutoGPTQ only accepts 2/3/4/8-bit quantization.")
|
|
128
|
+
|
|
129
|
+
require_version("optimum>=1.17.0", "To fix: pip install optimum>=1.17.0")
|
|
130
|
+
require_version("auto_gptq>=0.5.0", "To fix: pip install auto_gptq>=0.5.0")
|
|
131
|
+
from accelerate.utils import get_max_memory
|
|
132
|
+
|
|
133
|
+
if getattr(config, "model_type", None) == "chatglm":
|
|
134
|
+
raise ValueError("ChatGLM model is not supported yet.")
|
|
135
|
+
|
|
136
|
+
init_kwargs["quantization_config"] = GPTQConfig(
|
|
137
|
+
bits=quantization_args.export_quantization_bit,
|
|
138
|
+
dataset=_get_quantization_dataset(tokenizer, quantization_args),
|
|
139
|
+
)
|
|
140
|
+
init_kwargs["device_map"] = "auto"
|
|
141
|
+
init_kwargs["max_memory"] = get_max_memory()
|
|
142
|
+
logger.info("Quantizing model to {} bit with AutoGPTQ.".format(quantization_args.export_quantization_bit))
|
|
143
|
+
elif quantization_args.quantization_bit is not None: # on-the-fly
|
|
144
|
+
if quantization_args.quantization_method == QuantizationMethod.BITS_AND_BYTES.value:
|
|
145
|
+
if quantization_args.quantization_bit == 8:
|
|
146
|
+
require_version("bitsandbytes>=0.37.0", "To fix: pip install bitsandbytes>=0.37.0")
|
|
147
|
+
init_kwargs["quantization_config"] = BitsAndBytesConfig(load_in_8bit=True)
|
|
148
|
+
elif quantization_args.quantization_bit == 4:
|
|
149
|
+
require_version("bitsandbytes>=0.39.0", "To fix: pip install bitsandbytes>=0.39.0")
|
|
150
|
+
init_kwargs["quantization_config"] = BitsAndBytesConfig(
|
|
151
|
+
load_in_4bit=True,
|
|
152
|
+
bnb_4bit_compute_dtype=torch_dtype,
|
|
153
|
+
bnb_4bit_use_double_quant=quantization_args.double_quantization,
|
|
154
|
+
bnb_4bit_quant_type=quantization_args.quantization_type,
|
|
155
|
+
bnb_4bit_quant_storage=torch_dtype, # crucial for fsdp+qlora
|
|
156
|
+
)
|
|
157
|
+
else:
|
|
158
|
+
raise ValueError("Bitsandbytes only accepts 4-bit or 8-bit quantization.")
|
|
159
|
+
|
|
160
|
+
# Do not assign device map if:
|
|
161
|
+
# 1. deepspeed zero3 or fsdp (train)
|
|
162
|
+
# 2. auto quantization device map (inference)
|
|
163
|
+
if is_deepspeed_zero3_enabled() or is_fsdp_enabled() or quantization_args.quantization_device_map == "auto":
|
|
164
|
+
if quantization_args.quantization_bit != 4:
|
|
165
|
+
raise ValueError("Only 4-bit quantized model can use fsdp+qlora or auto device map.")
|
|
166
|
+
|
|
167
|
+
require_version("bitsandbytes>=0.43.0", "To fix: pip install bitsandbytes>=0.43.0")
|
|
168
|
+
else:
|
|
169
|
+
init_kwargs["device_map"] = {"": get_current_device()} # change auto device map for inference
|
|
170
|
+
|
|
171
|
+
logger.info("Quantizing model to {} bit with bitsandbytes.".format(quantization_args.quantization_bit))
|
|
172
|
+
elif quantization_args.quantization_method == QuantizationMethod.HQQ.value:
|
|
173
|
+
if quantization_args.quantization_bit not in [8, 6, 5, 4, 3, 2, 1]:
|
|
174
|
+
raise ValueError("HQQ only accepts 1/2/3/4/5/6/8-bit quantization.")
|
|
175
|
+
|
|
176
|
+
if is_deepspeed_zero3_enabled() or is_fsdp_enabled():
|
|
177
|
+
raise ValueError("HQQ quantization is incompatible with DeepSpeed ZeRO-3 or FSDP.")
|
|
178
|
+
|
|
179
|
+
require_version("hqq", "To fix: pip install hqq")
|
|
180
|
+
init_kwargs["quantization_config"] = HqqConfig(
|
|
181
|
+
nbits=quantization_args.quantization_bit, quant_zero=False, quant_scale=False, axis=0
|
|
182
|
+
) # use ATEN kernel (axis=0) for performance
|
|
183
|
+
logger.info("Quantizing model to {} bit with HQQ.".format(quantization_args.quantization_bit))
|
|
184
|
+
elif quantization_args.quantization_method == QuantizationMethod.EETQ.value:
|
|
185
|
+
if quantization_args.quantization_bit != 8:
|
|
186
|
+
raise ValueError("EETQ only accepts 8-bit quantization.")
|
|
187
|
+
|
|
188
|
+
if is_deepspeed_zero3_enabled() or is_fsdp_enabled():
|
|
189
|
+
raise ValueError("EETQ quantization is incompatible with DeepSpeed ZeRO-3 or FSDP.")
|
|
190
|
+
|
|
191
|
+
require_version("eetq", "To fix: pip install eetq")
|
|
192
|
+
init_kwargs["quantization_config"] = EetqConfig()
|
|
193
|
+
logger.info("Quantizing model to {} bit with EETQ.".format(quantization_args.quantization_bit))
|
|
@@ -0,0 +1,55 @@
|
|
|
1
|
+
from types import MethodType
|
|
2
|
+
from typing import Dict, Any
|
|
3
|
+
|
|
4
|
+
from transformers import AutoTokenizer, PreTrainedTokenizerBase, PreTrainedTokenizer
|
|
5
|
+
|
|
6
|
+
from distillflow.common import get_logger
|
|
7
|
+
from distillflow.model.args import ModelArgs
|
|
8
|
+
|
|
9
|
+
logger = get_logger(__name__)
|
|
10
|
+
|
|
11
|
+
def tokenizer_init_kwargs(model_args: ModelArgs) -> Dict[str, Any]:
|
|
12
|
+
r"""
|
|
13
|
+
Gets arguments to load config/tokenizer/model.
|
|
14
|
+
|
|
15
|
+
Note: including inplace operation of model_args.
|
|
16
|
+
"""
|
|
17
|
+
return {
|
|
18
|
+
"trust_remote_code": True,
|
|
19
|
+
"cache_dir": model_args.cache_dir,
|
|
20
|
+
"revision": model_args.model_revision,
|
|
21
|
+
"token": model_args.hf_hub_token,
|
|
22
|
+
}
|
|
23
|
+
|
|
24
|
+
def load_tokenizer(model_args: ModelArgs, template: str = None, padding_side='right') -> PreTrainedTokenizer:
|
|
25
|
+
try:
|
|
26
|
+
tokenizer = AutoTokenizer.from_pretrained(
|
|
27
|
+
model_args.model_name_or_path,
|
|
28
|
+
split_special_tokens=model_args.split_special_tokens,
|
|
29
|
+
padding_side=padding_side,
|
|
30
|
+
**tokenizer_init_kwargs(model_args),
|
|
31
|
+
)
|
|
32
|
+
except Exception as e:
|
|
33
|
+
raise OSError("Failed to load tokenizer.") from e
|
|
34
|
+
|
|
35
|
+
if model_args.new_special_tokens is not None:
|
|
36
|
+
print(model_args.new_special_tokens)
|
|
37
|
+
# exit()
|
|
38
|
+
num_added_tokens = tokenizer.add_special_tokens(
|
|
39
|
+
dict(additional_special_tokens=model_args.new_special_tokens.split(',')),
|
|
40
|
+
replace_additional_special_tokens=False,
|
|
41
|
+
)
|
|
42
|
+
logger.info("Add {} to special tokens.".format(model_args.new_special_tokens))
|
|
43
|
+
if num_added_tokens > 0 and not model_args.resize_vocab:
|
|
44
|
+
model_args.resize_vocab = True
|
|
45
|
+
logger.warning("New tokens have been added, changed `resize_vocab` to True.")
|
|
46
|
+
else:
|
|
47
|
+
tokenizer.pad_token = tokenizer.eos_token
|
|
48
|
+
|
|
49
|
+
if "PreTrainedTokenizerBase" not in str(tokenizer._pad.__func__):
|
|
50
|
+
tokenizer._pad = MethodType(PreTrainedTokenizerBase._pad, tokenizer)
|
|
51
|
+
|
|
52
|
+
if model_args.chat_template is not None:
|
|
53
|
+
tokenizer.chat_template = model_args.chat_template
|
|
54
|
+
|
|
55
|
+
return tokenizer
|
|
@@ -0,0 +1,92 @@
|
|
|
1
|
+
from typing import Any, Dict, Optional
|
|
2
|
+
|
|
3
|
+
import torch
|
|
4
|
+
|
|
5
|
+
from ..common import get_current_device, infer_optim_dtype
|
|
6
|
+
from ..common.logger import get_logger
|
|
7
|
+
|
|
8
|
+
from transformers import PretrainedConfig, PreTrainedModel
|
|
9
|
+
|
|
10
|
+
from .args import ModelArgs
|
|
11
|
+
|
|
12
|
+
logger = get_logger(__name__)
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def _get_unsloth_kwargs(
|
|
16
|
+
config: PretrainedConfig, model_name_or_path: str, model_args: ModelArgs
|
|
17
|
+
) -> Dict[str, Any]:
|
|
18
|
+
if model_args.infer_dtype != "auto":
|
|
19
|
+
torch_dtype = getattr(torch, model_args.infer_dtype)
|
|
20
|
+
else:
|
|
21
|
+
torch_dtype = infer_optim_dtype(model_dtype=getattr(config, "torch_dtype", None))
|
|
22
|
+
|
|
23
|
+
return {
|
|
24
|
+
"model_name": model_name_or_path,
|
|
25
|
+
"max_seq_length": 4096,
|
|
26
|
+
"dtype": torch_dtype,
|
|
27
|
+
"load_in_4bit": model_args.quantization_args.quantization_bit == 4,
|
|
28
|
+
"token": model_args.hf_hub_token,
|
|
29
|
+
"device_map": {"": get_current_device()},
|
|
30
|
+
"rope_scaling": getattr(config, "rope_scaling", None),
|
|
31
|
+
"fix_tokenizer": False,
|
|
32
|
+
"trust_remote_code": True,
|
|
33
|
+
"use_gradient_checkpointing": "unsloth",
|
|
34
|
+
}
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
def load_unsloth_pretrained_model(
|
|
38
|
+
config: PretrainedConfig, model_args: ModelArgs
|
|
39
|
+
) -> Optional[PreTrainedModel]:
|
|
40
|
+
r"""
|
|
41
|
+
Optionally loads pretrained model with unsloth. Used in training.
|
|
42
|
+
"""
|
|
43
|
+
from unsloth import FastLanguageModel
|
|
44
|
+
|
|
45
|
+
unsloth_kwargs = _get_unsloth_kwargs(config, model_args.model_name_or_path, model_args)
|
|
46
|
+
try:
|
|
47
|
+
model, _ = FastLanguageModel.from_pretrained(**unsloth_kwargs)
|
|
48
|
+
except NotImplementedError:
|
|
49
|
+
logger.warning("Unsloth does not support model type {}.".format(getattr(config, "model_type", None)))
|
|
50
|
+
model = None
|
|
51
|
+
model_args.use_unsloth = False
|
|
52
|
+
|
|
53
|
+
return model
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
def get_unsloth_peft_model(
|
|
57
|
+
model: "PreTrainedModel", model_args: "ModelArgs", peft_kwargs: Dict[str, Any]
|
|
58
|
+
) -> "PreTrainedModel":
|
|
59
|
+
r"""
|
|
60
|
+
Gets the peft model for the pretrained model with unsloth. Used in training.
|
|
61
|
+
"""
|
|
62
|
+
from unsloth import FastLanguageModel
|
|
63
|
+
|
|
64
|
+
unsloth_peft_kwargs = {
|
|
65
|
+
"model": model,
|
|
66
|
+
"max_seq_length": 4096,
|
|
67
|
+
"use_gradient_checkpointing": "unsloth",
|
|
68
|
+
}
|
|
69
|
+
return FastLanguageModel.get_peft_model(**peft_kwargs, **unsloth_peft_kwargs)
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
def load_unsloth_peft_model(
|
|
73
|
+
config: "PretrainedConfig", model_args: "ModelArgs", is_trainable: bool
|
|
74
|
+
) -> "PreTrainedModel":
|
|
75
|
+
r"""
|
|
76
|
+
Loads peft model with unsloth. Used in both training and inference.
|
|
77
|
+
"""
|
|
78
|
+
from unsloth import FastLanguageModel
|
|
79
|
+
|
|
80
|
+
unsloth_kwargs = _get_unsloth_kwargs(config, model_args.adapter_name_or_path[0], model_args)
|
|
81
|
+
try:
|
|
82
|
+
if not is_trainable:
|
|
83
|
+
unsloth_kwargs["use_gradient_checkpointing"] = False
|
|
84
|
+
|
|
85
|
+
model, _ = FastLanguageModel.from_pretrained(**unsloth_kwargs)
|
|
86
|
+
except NotImplementedError:
|
|
87
|
+
raise ValueError("Unsloth does not support model type {}.".format(getattr(config, "model_type", None)))
|
|
88
|
+
|
|
89
|
+
if not is_trainable:
|
|
90
|
+
FastLanguageModel.for_inference(model)
|
|
91
|
+
|
|
92
|
+
return model
|
|
@@ -0,0 +1,82 @@
|
|
|
1
|
+
from typing import List, Dict
|
|
2
|
+
|
|
3
|
+
import torch
|
|
4
|
+
|
|
5
|
+
class AdaptationLayer(torch.nn.Module):
|
|
6
|
+
def __init__(self, student_dim,
|
|
7
|
+
teacher_dim,
|
|
8
|
+
num_student_layers: int,
|
|
9
|
+
num_teacher_layers: int,
|
|
10
|
+
strategy="interpolate",
|
|
11
|
+
dtype=torch.bfloat16,
|
|
12
|
+
selection_indices: List[int]=None,
|
|
13
|
+
weights:List[List[int]] =None):
|
|
14
|
+
super().__init__()
|
|
15
|
+
self.projections = torch.nn.ModuleList([
|
|
16
|
+
torch.nn.Linear(student_dim, teacher_dim, dtype=dtype)
|
|
17
|
+
for _ in range(num_student_layers)
|
|
18
|
+
])
|
|
19
|
+
# self.layer_mapping = self.create_layer_mapping(num_student_layers, num_teacher_layers)
|
|
20
|
+
self.layer_mapping = self.map_teacher_to_student_layers(num_student_layers, num_teacher_layers, strategy, selection_indices, weights)
|
|
21
|
+
self.dtype = dtype
|
|
22
|
+
|
|
23
|
+
def map_teacher_to_student_layers(self, num_student_layers, num_teacher_layers, strategy="select",
|
|
24
|
+
selection_indices:List[int]=None, weights:List[List[int]]=None) -> {}:
|
|
25
|
+
"""
|
|
26
|
+
Maps teacher model layers to student model layers based on the specified strategy.
|
|
27
|
+
|
|
28
|
+
Args:
|
|
29
|
+
num_student_layers (int): Number of layers in the student model.
|
|
30
|
+
num_teacher_layers (int): Number of layers in the teacher model.
|
|
31
|
+
strategy (str): Layer mapping strategy ("direct", "select", "interpolate", "weighted").
|
|
32
|
+
selection_indices (list): Specific teacher layers to select (used when strategy="select").
|
|
33
|
+
weights (list of lists): Weights for combining teacher layers for each student layer
|
|
34
|
+
(used when strategy="weighted").
|
|
35
|
+
|
|
36
|
+
Returns:
|
|
37
|
+
list: List of mapping indices or weights from teacher layers to align with student layers.
|
|
38
|
+
"""
|
|
39
|
+
if strategy == "direct":
|
|
40
|
+
# Direct one-to-one mapping
|
|
41
|
+
return {
|
|
42
|
+
i: i
|
|
43
|
+
for i in range(num_student_layers)
|
|
44
|
+
}
|
|
45
|
+
|
|
46
|
+
elif strategy == "select":
|
|
47
|
+
# Use specific teacher layers for mapping
|
|
48
|
+
if selection_indices is None:
|
|
49
|
+
raise ValueError("selection_indices must be provided for 'select' strategy.")
|
|
50
|
+
if len(selection_indices) != num_student_layers:
|
|
51
|
+
raise ValueError("Number of selection_indices must match num_student_layers.")
|
|
52
|
+
return {
|
|
53
|
+
i: teacher_layer
|
|
54
|
+
for i, teacher_layer in enumerate(selection_indices)
|
|
55
|
+
}
|
|
56
|
+
|
|
57
|
+
elif strategy == "interpolate":
|
|
58
|
+
# Interpolate teacher layers to match student layers
|
|
59
|
+
return {
|
|
60
|
+
i: round(i * (num_teacher_layers - 1) / (num_student_layers - 1))
|
|
61
|
+
for i in range(num_student_layers)
|
|
62
|
+
}
|
|
63
|
+
elif strategy == "weighted":
|
|
64
|
+
# Weighted combination of teacher layers for each student layer
|
|
65
|
+
if weights is None:
|
|
66
|
+
raise ValueError("weights must be provided for 'weighted' strategy.")
|
|
67
|
+
if len(weights) != num_student_layers:
|
|
68
|
+
raise ValueError("Number of weight sets must match num_student_layers.")
|
|
69
|
+
return {
|
|
70
|
+
i: {j: weight for j, weight in enumerate(weight_set) if weight > 0}
|
|
71
|
+
for i, weight_set in enumerate(weights)
|
|
72
|
+
}
|
|
73
|
+
else:
|
|
74
|
+
raise ValueError(f"Unknown strategy: {strategy}")
|
|
75
|
+
|
|
76
|
+
def forward(self, student_hidden_states):
|
|
77
|
+
adapted_hidden_states = []
|
|
78
|
+
for i, hidden_state in enumerate(student_hidden_states):
|
|
79
|
+
if i >= len(self.projections):
|
|
80
|
+
break
|
|
81
|
+
adapted_hidden_states.append(self.projections[i](hidden_state.to(self.dtype)))
|
|
82
|
+
return adapted_hidden_states
|
|
File without changes
|
|
@@ -0,0 +1,56 @@
|
|
|
1
|
+
from typing import Literal, Optional, List, Dict
|
|
2
|
+
|
|
3
|
+
from pydantic import BaseModel, Field
|
|
4
|
+
from trl import SFTConfig
|
|
5
|
+
|
|
6
|
+
class DistillArgs(BaseModel):
|
|
7
|
+
sft_config: SFTConfig = Field(
|
|
8
|
+
description="SFT Config, with hyperparameters required for training"
|
|
9
|
+
)
|
|
10
|
+
type: Literal["logits", "layers", "attention", "fine-tune"] = Field(
|
|
11
|
+
default="logits",
|
|
12
|
+
description="Type of distillation to perform on the given student model with the given dataset (default: logits)"
|
|
13
|
+
)
|
|
14
|
+
max_seq_length: Optional[int] = Field (
|
|
15
|
+
default=4096,
|
|
16
|
+
description="Maximum sequence length to use during training (default: 4096)",
|
|
17
|
+
examples=[1024, 2048, 4096]
|
|
18
|
+
)
|
|
19
|
+
dataset_text_field: Optional[str] = Field(
|
|
20
|
+
default="text",
|
|
21
|
+
description="The key to which the data is mapped",
|
|
22
|
+
)
|
|
23
|
+
temperature: Optional[float] = Field(
|
|
24
|
+
default=0.5,
|
|
25
|
+
description="Temperature"
|
|
26
|
+
)
|
|
27
|
+
alpha: Optional[float] = Field(
|
|
28
|
+
default=2.0,
|
|
29
|
+
description="alpha"
|
|
30
|
+
)
|
|
31
|
+
resume_from_checkpoint: Optional[str] = Field(
|
|
32
|
+
default=None,
|
|
33
|
+
description="Training checkpoint folder path to resume the training from which to resume the training"
|
|
34
|
+
)
|
|
35
|
+
strategy : Literal["direct", "select", "interpolate", "weighted"] = Field(
|
|
36
|
+
default = "interpolate",
|
|
37
|
+
description="Strategy to select when mapping the teacher and student layers/attention map"
|
|
38
|
+
)
|
|
39
|
+
selection_indices: Optional[List[int]] = Field(
|
|
40
|
+
default=None,
|
|
41
|
+
description="If selected strategy `select`, provides mapping between student and teacher model layer/attention map",
|
|
42
|
+
examples=[[1,2,0]] # mapping 0th student layer to 1st layer of teacher and so on
|
|
43
|
+
)
|
|
44
|
+
weights: Optional[List[List[int]]] = Field(
|
|
45
|
+
default=None,
|
|
46
|
+
description="If selected strategy `weighted`, provides the weights for each of the teacher layers to be used when computing student layer",
|
|
47
|
+
examples=[[[100,10,40,50], [10,5,50,25]]] # 0th of student is computed with teacher layer weights 100 for 0th layer, 10 for 1st layer and so on
|
|
48
|
+
)
|
|
49
|
+
|
|
50
|
+
remove_unused_columns: Optional[bool] = Field(
|
|
51
|
+
default=False,
|
|
52
|
+
)
|
|
53
|
+
|
|
54
|
+
model_config = {
|
|
55
|
+
"extra": "forbid"
|
|
56
|
+
}
|
|
@@ -0,0 +1,113 @@
|
|
|
1
|
+
from typing import Union, Optional
|
|
2
|
+
|
|
3
|
+
import torch
|
|
4
|
+
from accelerate import Accelerator
|
|
5
|
+
from datasets import IterableDataset
|
|
6
|
+
import torch.nn.functional as F
|
|
7
|
+
from transformers import PreTrainedModel, PreTrainedTokenizerBase
|
|
8
|
+
from trl import SFTTrainer
|
|
9
|
+
|
|
10
|
+
from distillflow.common import get_current_device
|
|
11
|
+
from distillflow.datasets.loader import DatasetModule
|
|
12
|
+
from distillflow.trainer.AdaptationLayer import AdaptationLayer
|
|
13
|
+
from distillflow.trainer.args import DistillArgs
|
|
14
|
+
|
|
15
|
+
class AttentionTrainer(SFTTrainer):
|
|
16
|
+
def __init__(self,
|
|
17
|
+
accelerator: Accelerator,
|
|
18
|
+
distill_args: DistillArgs,
|
|
19
|
+
teacher_model: PreTrainedModel,
|
|
20
|
+
model: PreTrainedModel,
|
|
21
|
+
dataset_module: DatasetModule,
|
|
22
|
+
tokenizer: PreTrainedTokenizerBase
|
|
23
|
+
):
|
|
24
|
+
self.teacher_model = teacher_model
|
|
25
|
+
self.distill_args = distill_args
|
|
26
|
+
train_dataset = dataset_module["train_dataset"]
|
|
27
|
+
eval_dataset = dataset_module["eval_dataset"]
|
|
28
|
+
self.device = get_current_device()
|
|
29
|
+
self.adaptation_layer = AdaptationLayer(
|
|
30
|
+
model.config.hidden_size,
|
|
31
|
+
teacher_model.config.hidden_size,
|
|
32
|
+
model.config.num_hidden_layers,
|
|
33
|
+
teacher_model.config.num_hidden_layers,
|
|
34
|
+
dtype=torch.float16,
|
|
35
|
+
strategy=distill_args.strategy,
|
|
36
|
+
selection_indices=distill_args.selection_indices,
|
|
37
|
+
weights=distill_args.weights
|
|
38
|
+
).to(self.device)
|
|
39
|
+
|
|
40
|
+
if isinstance(train_dataset, IterableDataset) and distill_args.sft_config.max_steps == -1:
|
|
41
|
+
raise ValueError("max steps should be specified when using dataset with streaming mode enabled.")
|
|
42
|
+
|
|
43
|
+
distill_args.sft_config.max_length = distill_args.max_seq_length
|
|
44
|
+
distill_args.sft_config.dataset_text_field = distill_args.dataset_text_field
|
|
45
|
+
|
|
46
|
+
super().__init__(model=model, args=distill_args.sft_config, train_dataset=train_dataset,
|
|
47
|
+
eval_dataset=eval_dataset, processing_class=tokenizer)
|
|
48
|
+
|
|
49
|
+
def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=None):
|
|
50
|
+
# Forward pass for the student model
|
|
51
|
+
student_outputs = model(**inputs, output_attentions=True)
|
|
52
|
+
|
|
53
|
+
self.teacher_model = self.teacher_model.to(self.device) if self.device.type == "mps" else self.teacher_model
|
|
54
|
+
|
|
55
|
+
teacher_model = self.teacher_model.module if hasattr(self.teacher_model, 'module') else self.teacher_model
|
|
56
|
+
|
|
57
|
+
# Forward pass for the teacher model
|
|
58
|
+
with torch.no_grad():
|
|
59
|
+
teacher_outputs = teacher_model(**inputs, output_attentions=True)
|
|
60
|
+
|
|
61
|
+
# Primary task loss (e.g., cross-entropy loss)
|
|
62
|
+
loss = student_outputs.loss
|
|
63
|
+
|
|
64
|
+
# Attention-based distillation loss
|
|
65
|
+
attention_loss = self.compute_attention_loss(
|
|
66
|
+
teacher_outputs.attentions,
|
|
67
|
+
student_outputs.attentions
|
|
68
|
+
)
|
|
69
|
+
|
|
70
|
+
# Combine losses
|
|
71
|
+
total_loss = ((1 - self.distill_args.alpha) * loss + self.distill_args.alpha * attention_loss) / self.args.gradient_accumulation_steps
|
|
72
|
+
|
|
73
|
+
return (total_loss, student_outputs) if return_outputs else total_loss
|
|
74
|
+
|
|
75
|
+
def compute_attention_loss(self, teacher_attentions, student_attentions):
|
|
76
|
+
"""
|
|
77
|
+
Compute attention-based distillation loss.
|
|
78
|
+
|
|
79
|
+
Args:
|
|
80
|
+
teacher_attentions: List of teacher model attention maps.
|
|
81
|
+
student_attentions: List of student model attention maps.
|
|
82
|
+
|
|
83
|
+
Returns:
|
|
84
|
+
Total attention distillation loss.
|
|
85
|
+
"""
|
|
86
|
+
loss = 0.0
|
|
87
|
+
num_layers = len(student_attentions)
|
|
88
|
+
|
|
89
|
+
self.adaptation_layer = self.adaptation_layer.to(self.device)
|
|
90
|
+
|
|
91
|
+
for student_idx, teacher_idx in self.adaptation_layer.layer_mapping.items():
|
|
92
|
+
if self.distill_args.strategy == "weighted":
|
|
93
|
+
teacher_attention = torch.zeros_like(teacher_attentions[0])
|
|
94
|
+
for idx, weight in teacher_idx.items():
|
|
95
|
+
teacher_attention = weight * teacher_attentions[idx]
|
|
96
|
+
else:
|
|
97
|
+
teacher_attention = teacher_attentions[teacher_idx]
|
|
98
|
+
|
|
99
|
+
student_attention = student_attentions[student_idx]
|
|
100
|
+
# Align dimensions if needed
|
|
101
|
+
if teacher_attention.size() != student_attention.size():
|
|
102
|
+
teacher_attention = teacher_attention.mean(dim=1, keepdim=True).expand(-1, student_attention.size(1), -1, -1)
|
|
103
|
+
teacher_attention = F.interpolate(
|
|
104
|
+
teacher_attention,
|
|
105
|
+
size=student_attention.size()[-2:], # Resize spatial dimensions
|
|
106
|
+
mode="bilinear",
|
|
107
|
+
align_corners=False,
|
|
108
|
+
)
|
|
109
|
+
|
|
110
|
+
# MSE Loss for attention maps
|
|
111
|
+
loss += F.mse_loss(student_attention, teacher_attention)
|
|
112
|
+
|
|
113
|
+
return loss / num_layers
|
|
@@ -0,0 +1,51 @@
|
|
|
1
|
+
from accelerate import Accelerator
|
|
2
|
+
from datasets import IterableDataset
|
|
3
|
+
from transformers import PreTrainedModel, PreTrainedTokenizerBase
|
|
4
|
+
from trl import SFTTrainer
|
|
5
|
+
import torch
|
|
6
|
+
import torch.nn.functional as F
|
|
7
|
+
|
|
8
|
+
from .args import DistillArgs
|
|
9
|
+
from ..common import get_current_device
|
|
10
|
+
from ..datasets.loader import DatasetModule
|
|
11
|
+
|
|
12
|
+
class FineTuning(SFTTrainer):
|
|
13
|
+
def __init__(self,
|
|
14
|
+
accelerator: Accelerator,
|
|
15
|
+
distill_args: DistillArgs,
|
|
16
|
+
teacher_model: PreTrainedModel,
|
|
17
|
+
model: PreTrainedModel,
|
|
18
|
+
dataset_module: DatasetModule,
|
|
19
|
+
tokenizer: PreTrainedTokenizerBase
|
|
20
|
+
):
|
|
21
|
+
self.accelerator = accelerator
|
|
22
|
+
self.distill_args = distill_args
|
|
23
|
+
train_dataset = dataset_module["train_dataset"]
|
|
24
|
+
eval_dataset = dataset_module["eval_dataset"]
|
|
25
|
+
self.device = get_current_device()
|
|
26
|
+
if self.device.type == 'mps':
|
|
27
|
+
# Explicitly place the models on device since accelerate prepare does not work on MPS.
|
|
28
|
+
model = model.to(self.device)
|
|
29
|
+
|
|
30
|
+
if isinstance(train_dataset, IterableDataset) and distill_args.sft_config.max_steps == -1:
|
|
31
|
+
raise ValueError("max steps should be specified when using dataset with streaming mode enabled.")
|
|
32
|
+
|
|
33
|
+
distill_args.sft_config.max_length = distill_args.max_seq_length
|
|
34
|
+
distill_args.sft_config.dataset_text_field = distill_args.dataset_text_field
|
|
35
|
+
|
|
36
|
+
super().__init__(model=model, args=distill_args.sft_config, train_dataset=train_dataset,
|
|
37
|
+
eval_dataset=eval_dataset, processing_class=tokenizer)
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def output(self, model, inputs, no_grad):
|
|
41
|
+
if no_grad:
|
|
42
|
+
with torch.no_grad():
|
|
43
|
+
return model(**inputs)
|
|
44
|
+
else:
|
|
45
|
+
return model(**inputs)
|
|
46
|
+
|
|
47
|
+
def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=None):
|
|
48
|
+
student_model = model.module if hasattr(model, 'module') else model
|
|
49
|
+
student_outputs = self.output(student_model, inputs, False)
|
|
50
|
+
|
|
51
|
+
return student_outputs.loss
|