大模型蒸馏算法代码
·
import os
import torch
import torch.nn as nn
import torch.nn.functional as F
from dataclasses import dataclass, field
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
Trainer,
TrainingArguments,
DataCollatorForLanguageModeling,
HfArgumentParser,
)
from datasets import load_dataset
from typing import Optional
from functools import partial # <--- 新增导入
# ===================================================================================
# 1. 参数定义 (无变化)
# ===================================================================================
@dataclass
class ModelArguments:
teacher_model_name_or_path: str = field(
metadata={"help": "教师模型的路径或Hugging Face Hub上的名称。"}
)
student_model_name_or_path: str = field(
metadata={"help": "学生模型的路径或Hugging Face Hub上的名称。"}
)
@dataclass
class DataArguments:
dataset_name: str = field(
metadata={"help": "Hugging Face Hub上的数据集名称,或指向本地数据文件(如.jsonl, .csv)的路径。"}
)
max_seq_length: int = field(
default=512, metadata={"help": "分词后的最大序列长度。"}
)
@dataclass
class DistillationTrainingArguments(TrainingArguments):
output_dir: str = field(default="./distilled-model", metadata={"help": "模型输出和检查点的目录。"})
distillation_alpha: float = field(
default=0.5, metadata={"help": "蒸馏损失和学生损失之间的权重。"}
)
distillation_temperature: float = field(
default=2.0, metadata={"help": "蒸馏中的温度参数。"}
)
# ===================================================================================
# 2. 自定义蒸馏 Trainer (无变化)
# ===================================================================================
class DistillationTrainer(Trainer):
def __init__(self, *args, teacher_model=None, **kwargs):
super().__init__(*args, **kwargs)
self.teacher = teacher_model
if self.teacher:
self.teacher.to(self.args.device)
self.teacher.eval()
def compute_loss(self, model, inputs, return_outputs=False, **kwargs):
outputs_student = model(**inputs)
loss_ce = outputs_student.loss
logits_student = outputs_student.logits
with torch.no_grad():
outputs_teacher = self.teacher(**inputs)
logits_teacher = outputs_teacher.logits
alpha = self.args.distillation_alpha
temperature = self.args.distillation_temperature
loss_kd = nn.KLDivLoss(reduction="batchmean")(
F.log_softmax(logits_student / temperature, dim=-1),
F.softmax(logits_teacher / temperature, dim=-1)
) * (temperature ** 2)
loss = alpha * loss_ce + (1 - alpha) * loss_kd
return (loss, outputs_student) if return_outputs else loss
# ===================================================================================
# 3. Alpaca 数据集预处理函数 (***核心修改处***)
# ===================================================================================
PROMPT_DICT = {
"prompt_input": (
"Below is an instruction that describes a task, paired with an input that provides further context. "
"Write a response that appropriately completes the request.\n\n"
"### Instruction:\n{instruction}\n\n### Input:\n{input}\n\n### Response:"
),
"prompt_no_input": (
"Below is an instruction that describes a task. "
"Write a response that appropriately completes the request.\n\n"
"### Instruction:\n{instruction}\n\n### Response:"
),
}
def preprocess_function(examples, tokenizer, max_length):
"""
对Alpaca数据集进行预处理和分词。
这个版本使用了静态填充(padding="max_length")来避免动态填充的错误。
"""
# 1. 格式化prompt
prompts = []
for instruction, input_text in zip(examples['instruction'], examples['input']):
if input_text and input_text.strip() != "":
prompts.append(PROMPT_DICT["prompt_input"].format(instruction=instruction, input=input_text))
else:
prompts.append(PROMPT_DICT["prompt_no_input"].format(instruction=instruction))
# 2. 将prompt和response拼接
full_texts = [prompt + " " + output for prompt, output in zip(prompts, examples['output'])]
# 3. 对完整文本进行分词,并直接填充到max_length
model_inputs = tokenizer(
full_texts,
padding="max_length", # <-- 修改点: 静态填充
truncation=True,
max_length=max_length
)
# 4. 为了创建labels,我们需要知道prompt部分的长度
# 我们只对prompt进行分词(不填充),以获取其真实长度
prompt_tokens = tokenizer(prompts, padding=False, truncation=True, max_length=max_length)
# 5. 创建labels,并mask掉prompt部分
labels = [row.copy() for row in model_inputs["input_ids"]] # 深度复制
for i in range(len(labels)):
prompt_len = len(prompt_tokens['input_ids'][i])
labels[i][:prompt_len] = [-100] * prompt_len
# 另外一个细节:在填充区域,label也应该是-100
# tokenizer.pad() 会用 pad_token_id 填充,我们需要替换它
pad_token_id = tokenizer.pad_token_id
for j in range(len(labels[i])):
if model_inputs["input_ids"][i][j] == pad_token_id:
labels[i][j] = -100
model_inputs["labels"] = labels
return model_inputs
# ===================================================================================
# 主执行函数
# ===================================================================================
def main():
parser = HfArgumentParser((ModelArguments, DataArguments, DistillationTrainingArguments))
model_args, data_args, training_args = parser.parse_args_into_dataclasses()
# --- 加载模型和分词器 (无变化) ---
print("加载教师模型...")
teacher_model = AutoModelForCausalLM.from_pretrained(model_args.teacher_model_name_or_path)
print("加载学生模型...")
student_model = AutoModelForCausalLM.from_pretrained(model_args.student_model_name_or_path)
print("加载分词器...")
tokenizer = AutoTokenizer.from_pretrained(
model_args.student_model_name_or_path,
padding_side="right",
use_fast=False, # 使用慢速分词器通常更稳定
)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
student_model.config.pad_token_id = tokenizer.eos_token_id
teacher_model.config.pad_token_id = tokenizer.eos_token_id
# --- 加载和预处理数据集 (***.map()调用方式已修改***) ---
print("加载数据集...")
if os.path.exists(data_args.dataset_name):
print(f"从本地路径加载数据集: {data_args.dataset_name}")
raw_dataset = load_dataset("json", data_files={"train": data_args.dataset_name})
dataset = raw_dataset["train"]
else:
print(f"从 Hugging Face Hub 加载数据集: {data_args.dataset_name}")
dataset = load_dataset(data_args.dataset_name, split="train")
if len(dataset) > 1000:
dataset = dataset.select(range(1000))
print("预处理数据集...")
# 使用 functools.partial 将固定的参数(tokenizer, max_length)传入预处理函数
# 这是比 lambda 更好、更清晰的方式
preprocess_with_args = partial(
preprocess_function,
tokenizer=tokenizer,
max_length=data_args.max_seq_length
)
processed_dataset = dataset.map(
preprocess_with_args,
batched=True,
remove_columns=dataset.column_names,
desc="Running tokenizer on dataset",
)
# --- 初始化 Trainer (无变化) ---
print("初始化 Trainer...")
trainer = DistillationTrainer(
model=student_model,
teacher_model=teacher_model,
args=training_args,
train_dataset=processed_dataset,
tokenizer=tokenizer,
data_collator=DataCollatorForLanguageModeling(tokenizer, mlm=False),
)
# --- 开始训练和保存 (无变化) ---
print("开始蒸馏训练...")
train_result = trainer.train()
print("训练完成,保存模型...")
trainer.save_model()
tokenizer.save_pretrained(training_args.output_dir)
trainer.log_metrics("train", train_result.metrics)
trainer.save_metrics("train", train_result.metrics)
trainer.save_state()
if __name__ == "__main__":
main()
更多推荐



所有评论(0)