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()

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐