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,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