You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Llama 3.2 1B分层师生优化蒸馏模型不收敛,求排查方案

问题:Llama 3.2 1B分层师生蒸馏模型无法收敛

我尝试复现一篇论文中的分层师生蒸馏实验,针对Llama 3.2 1B模型,采用L2范数独立优化每个Transformer层,但模型无法收敛。已尝试调整优化器超参数(权重衰减、学习率)、更换自注意力权重初始化方式、增减最后一层归一化,问题仍未解决。相关实现代码如下:

import torch
import torch.nn as nn
import torch.optim as optim
from transformers import AutoModelForCausalLM, AutoTokenizer, LlamaConfig
from datasets import load_dataset
from lm_eval.models.huggingface import HFLM
from lm_eval import simple_evaluate

# ------------------------------------------------------------
# 1. 加载预训练教师模型与分词器
# ------------------------------------------------------------
model_name = "meta-llama/Llama-3.2-1B"  # 替换为实际模型名称
teacher_model = AutoModelForCausalLM.from_pretrained(model_name, attn_implementation="flash_attention_2")
tokenizer = AutoTokenizer.from_pretrained(model_name)

# 移动到GPU(如果可用)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
teacher_model.to(device).to(torch.bfloat16)
teacher_model.eval()

# ------------------------------------------------------------
# 2. 创建同配置的随机初始化学生模型
# ------------------------------------------------------------
config = teacher_model.config
student_model = AutoModelForCausalLM.from_config(config, attn_implementation="flash_attention_2")
student_model.to(device).to(torch.bfloat16)

# ------------------------------------------------------------
# 3. 复制教师模型权重,跳过自注意力q/k/v投影层
# ------------------------------------------------------------
with torch.no_grad():
    for (orig_name, orig_param), (rand_name, rand_param) in zip(teacher_model.named_parameters(), student_model.named_parameters()):
        if any(attn_name in orig_name for attn_name in ["q_proj", "k_proj", "v_proj"]):
            # 跳过注意力投影权重复制
            continue
        else:
            rand_param.data.copy_(orig_param.data)

# ------------------------------------------------------------
# 4. 准备数据集与数据加载器
# ------------------------------------------------------------
streaming_dataset = load_dataset("HuggingFaceFW/fineweb", split="train", streaming=True)

def tokenize_streaming_data(example):
    return tokenizer(example["text"], truncation=True, padding="max_length", max_length=1024)

streaming_dataset = streaming_dataset.map(tokenize_streaming_data)
dataloader = torch.utils.data.DataLoader(streaming_dataset, batch_size=8)

# ------------------------------------------------------------
# 5. 定义语言模型评估函数
# ------------------------------------------------------------
def lm_eval(student_model, tokenizer, device):

    results = simple_evaluate(
        model=HFLM(pretrained=student_model, tokenizer=tokenizer, backend="causal"),
        tasks=["arc_easy"],
        num_fewshot=0,
        device=device,
        log_samples=False,

    )

    print(f"步骤 {step} 评估结果:")
    for alias, result in results["results"].items():
        acc = result["acc,none"]
        acc_norm = result["acc_norm,none"]
        print(f"{alias}: {acc} (归一化准确率: {acc_norm})")

# ------------------------------------------------------------
# 6. 定义损失函数与优化器
# ------------------------------------------------------------
# 层输出间的L2损失
criterion = lambda x, y: torch.norm(x - y, p=2, dim=(-1,)).mean()
optimizer = optim.AdamW(student_model.parameters(), lr=1e-4)
tokenizer.pad_token = tokenizer.eos_token

student_model.train()
teacher_model.eval()

# ------------------------------------------------------------
# 7. 训练循环,每1000步进行一次评估
# ------------------------------------------------------------
step = 0
for batch in dataloader:
    step += 1
    input_ids = torch.stack(batch["input_ids"]).to(device).t()  # 形状 [seq_len, batch]
    attention_mask = torch.stack(batch["attention_mask"]).to(device).t()

    seq_len = input_ids.size(1)  # 每条序列的token数
    position_ids = torch.arange(seq_len, dtype=torch.long, device=device)
    position_ids = position_ids.unsqueeze(0).expand(input_ids.size(0), seq_len)  # [batch_size, seq_len]

    kwargs = {"attention_mask": attention_mask, "position_ids": position_ids}

    # 获取教师模型输出
    with torch.no_grad():
        output = teacher_model(input_ids, output_hidden_states=True, **kwargs)
        output = output["hidden_states"]

    # 计算学生模型与教师模型的层损失
    loss = 0.0
    # for i, layer in enumerate(student_model.model.layers):
    #     loss += criterion(layer(output[i], **kwargs)[0], output[i+1])
    for i, (student_layer, teacher_layer) in enumerate(zip(student_model.model.layers, teacher_model.model.layers)):
        loss += criterion(student_layer.self_attn(output[i], **kwargs)[0], 
                          teacher_layer.self_attn(output[i], **kwargs)[0])
    loss /= len(student_model.model.layers)

    optimizer.zero_grad()
    loss.backward()
    optimizer.step()

    if step % 100 == 0:
        print(f"步骤 {step} - 损失值: {loss.item()}")

    # 每1000步执行一次评估
    if step % 1000 == 0:
        student_model.eval()
        lm_eval(student_model, tokenizer, device)
        student_model.train()

内容的提问来源于stack exchange,提问作者Makao

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.15 11:22:04