大模型之OPD在线策略蒸馏训练实战篇

1、方案

步骤1:在线采样(On-Policy Sampling)

目的:用当前策略生成训练数据

过程

  • 输入一批prompts,使用当前Student模型生成完整回答

  • 记录生成序列、prompt长度,并计算当前策略下生成部分的log概率(student_old_logp

  • Teacher模型对同一序列计算生成部分的log概率(teacher_logp

  • 所有操作均在torch.no_grad()下执行,不产生梯度

关键点:采样用的是当前策略(最新模型),保证数据新鲜度

步骤2:梯度前向传播

目的:为反向传播构建计算图

过程

  • 将步骤1生成的完整序列再次输入Student模型

  • 计算生成部分的log概率(student_new_logp

  • 此过程开启梯度记录(student_train.train()

说明:第二次前向传播与第一次数值相同,但构建了计算图,支持后续梯度回传

步骤3:Loss计算

目的:量化Student与Teacher的差距

过程

  • 计算student_new_logpteacher_logp的KL散度

  • 目标:让Student的输出分布向Teacher靠拢

  • 将loss除以梯度累积步数,用于后续梯度累积

核心思想:Student模仿Teacher的token-level概率分布

步骤4:反向传播与更新

目的:优化Student模型参数

过程

  • 执行accelerator.backward(loss)计算梯度

  • 累积到GRAD_ACCUM_STEPS步后,执行梯度裁剪并更新参数

  • 优化器更新后,模型从旧权重W_old变为新权重W_new

关键点:梯度累积模拟更大batch size,梯度裁剪防止梯度爆炸

2、代码实现

import os
import json
import random
import copy
import torch
import torch.nn.functional as F
import warnings
from tqdm import tqdm
from datetime import datetime
from transformers import (
    AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig, GenerationConfig
)
from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training
from accelerate import Accelerator

warnings.filterwarnings("ignore")

# ===================== 路径配置 =====================
STUDENT_BASE_PATH = "/root/autodl-tmp/models/Qwen2.5-3B-Instruct"
TEACHER_MED_PATH = "/root/autodl-tmp/codes/sft/merged_sft_qwen7b_med"
TRAIN_DATA_PATH = "/root/autodl-tmp/datas/ppo/ppo_train_simple.jsonl"
VAL_DATA_PATH = "/root/autodl-tmp/datas/ppo/ppo_val_simple.jsonl"
SAVE_LORA_DIR = "./opd_3b_med_lora"
BEST_LORA_PATH = os.path.join(SAVE_LORA_DIR, "best_val_lora")
os.makedirs(SAVE_LORA_DIR, exist_ok=True)
LOG_FILE = os.path.join(SAVE_LORA_DIR, "train_log.txt")

# ===================== 超参数 =====================
LR = 2e-5
EPOCHS = 6
CLIP_EPS = 0.2
TEMPERATURE = 1.0
MAX_NEW_TOKENS = 768
GRAD_CLIP_NORM = 1.0
SAVE_STEP_INTERVAL = 200
VAL_SAMPLE_NUM = 200
BATCH_SIZE = 6
GRAD_ACCUM_STEPS = 2

# LoRA 配置
lora_config = LoraConfig(
    r=16,
    lora_alpha=32,
    target_modules=["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"],
    lora_dropout=0.05,
    bias="none",
    task_type="CAUSAL_LM"
)

# 4bit量化 + FP16混合精度
bnb_config = BitsAndBytesConfig(
    load_in_4bit=True,
    bnb_4bit_use_double_quant=True,
    bnb_4bit_quant_type="nf4",
    bnb_4bit_compute_dtype=torch.bfloat16
)

# 开启fp16加速
accelerator = Accelerator(mixed_precision="fp16")
device_train = accelerator.device
device_teacher = torch.device("cuda:1")
print(f"Student训练模型:{device_train}")
print(f"Teacher7B模型放置:{device_teacher}")

# ===================== Tokenizer =====================
tokenizer = AutoTokenizer.from_pretrained(TEACHER_MED_PATH, trust_remote_code=True)
tokenizer.pad_token = tokenizer.eos_token
tokenizer.padding_side = "left"

# ===================== 加载模型 =====================
print("加载 Student Qwen2.5-3B -> GPU0")
student_train = AutoModelForCausalLM.from_pretrained(
    STUDENT_BASE_PATH,
    quantization_config=bnb_config,
    device_map=device_train,
    torch_dtype=torch.bfloat16,
    trust_remote_code=True
)
student_train = prepare_model_for_kbit_training(student_train)
student_train.config.use_cache = False
student_train = get_peft_model(student_train, lora_config)
student_train.print_trainable_parameters()

print("加载 Teacher 7B医疗模型 -> GPU1")
teacher_model = AutoModelForCausalLM.from_pretrained(
    TEACHER_MED_PATH,
    quantization_config=bnb_config,
    device_map=device_teacher,
    torch_dtype=torch.bfloat16,
    trust_remote_code=True
).eval()
teacher_model.config.use_cache = False
for p in teacher_model.parameters():
    p.requires_grad = False


# ===================== 数据集 =====================
def load_prompts(jsonl_path):
    prompts = []
    with open(jsonl_path, "r", encoding="utf-8") as f:
        for line in f:
            line = line.strip()
            if not line:
                continue
            d = json.loads(line)
            prompts.append(d["prompt"])
    return prompts


train_prompts = load_prompts(TRAIN_DATA_PATH)
val_prompts_all = load_prompts(VAL_DATA_PATH)
print(f"训练集:{len(train_prompts)} 条 | 验证集:{len(val_prompts_all)} 条")


def chunk_list(lst, chunk_size):
    for i in range(0, len(lst), chunk_size):
        yield lst[i:i + chunk_size]


# ===================== 日志写入工具 =====================
def write_log(content: str):
    timestamp = datetime.now().strftime("%Y-%m-%d %H:%M:%S")
    log_line = f"[{timestamp}] {content}\n"
    with open(LOG_FILE, "a", encoding="utf-8") as f:
        f.write(log_line)


# ===================== 核心函数 =====================
@torch.no_grad()
def generate_and_compute_logp(model, prompt_batch):
    """
    用当前策略生成序列并计算logp(完全无梯度)
    返回:完整序列、prompt长度、生成部分的logp
    """
    model.eval()
    inputs = tokenizer(prompt_batch, return_tensors="pt", padding=True).to(device_train)
    prompt_len = inputs["input_ids"].size(1)

    gen_cfg = GenerationConfig(
        max_new_tokens=MAX_NEW_TOKENS,
        temperature=TEMPERATURE,
        do_sample=True,
        pad_token_id=tokenizer.eos_token_id,
        eos_token_id=tokenizer.eos_token_id,
        num_beams=1,
        use_cache=True
    )
    full_seq = model.generate(**inputs, generation_config=gen_cfg)

    # 提取生成的部分
    generated_ids = full_seq[:, prompt_len:]

    # 计算完整序列的logp
    out = model(full_seq)
    log_soft = F.log_softmax(out.logits, dim=-1)

    # 只取生成部分的logp
    old_logp = torch.gather(
        log_soft[:, prompt_len:, :],
        dim=-1,
        index=generated_ids.unsqueeze(-1)
    ).squeeze(-1)

    return full_seq, prompt_len, old_logp


@torch.no_grad()
def get_teacher_logprob(model, seq_ids, target_device):
    """计算teacher模型的log概率"""
    seq_ids = seq_ids.to(target_device)
    out = model(seq_ids)
    log_soft = F.log_softmax(out.logits, dim=-1)
    teacher_logp = torch.gather(log_soft, dim=-1, index=seq_ids.unsqueeze(-1)).squeeze(-1)
    return teacher_logp.to(device_train)


def compute_opd_loss(student_logp, teacher_logp):
    """
    计算OPD loss:直接最小化student和teacher的KL散度
    - student_logp: 当前策略在生成部分的logp(有梯度)
    - teacher_logp: teacher模型在生成部分的logp(无梯度)
    """
    # 维度对齐
    min_len = min(student_logp.size(1), teacher_logp.size(1))
    s_logp = student_logp[:, :min_len]
    t_logp = teacher_logp[:, :min_len]

    # 计算KL散度:KL(student || teacher)
    # 注意:这里student是当前策略,teacher是目标分布
    kl_loss = F.kl_div(
        F.log_softmax(s_logp, dim=-1),
        F.softmax(t_logp, dim=-1),
        reduction='batchmean'
    )

    return kl_loss


@torch.no_grad()
def run_validation(model, epoch, global_step):
    """验证:计算student和teacher的KL散度"""
    model.eval()
    val_sample = random.sample(val_prompts_all, min(VAL_SAMPLE_NUM, len(val_prompts_all)))

    total_kl = 0.0
    total_ce = 0.0
    count = 0

    pbar_val = tqdm(chunk_list(val_sample, BATCH_SIZE), desc="Validating", leave=False)
    for batch in pbar_val:
        full_seq, prompt_len, student_logp = generate_and_compute_logp(model, batch)
        teacher_logp_full = get_teacher_logprob(teacher_model, full_seq, device_teacher)

        # 提取teacher的生成部分
        teacher_logp = teacher_logp_full[:, prompt_len:]

        # 对齐维度
        min_len = min(student_logp.size(1), teacher_logp.size(1))
        student_logp_aligned = student_logp[:, :min_len]
        teacher_logp_aligned = teacher_logp[:, :min_len]

        # KL散度:teacher相对于student
        kl_div = (teacher_logp_aligned - student_logp_aligned).mean().item()

        # Cross-Entropy:student拟合teacher
        ce_loss = -teacher_logp_aligned.mean().item()

        batch_cnt = len(batch)
        total_kl += kl_div * batch_cnt
        total_ce += ce_loss * batch_cnt
        count += batch_cnt

        pbar_val.set_postfix({
            "KL": f"{total_kl / count:.4f}",
            "CE": f"{total_ce / count:.4f}"
        })

    avg_kl = total_kl / count if count > 0 else 0.0
    avg_ce = total_ce / count if count > 0 else 0.0
    write_log(f"【VAL】epoch:{epoch}, step:{global_step}, KL={avg_kl:.6f}, CE={avg_ce:.6f}")
    return avg_kl, avg_ce


# ===================== 优化器 =====================
optimizer = torch.optim.AdamW(student_train.parameters(), lr=LR)
student_train, optimizer = accelerator.prepare(student_train, optimizer)

global_step = 0
best_val_kl = float("inf")
write_log("========== Training Start ==========")

# ===================== 标准在线OPD主训练循环 =====================
for epoch in range(EPOCHS):
    print(f"\n===== Epoch {epoch}/{EPOCHS} =====")
    write_log(f"========== Epoch {epoch} Start ==========")

    # 打乱训练数据
    random.shuffle(train_prompts)
    train_batches = list(chunk_list(train_prompts, BATCH_SIZE))
    train_pbar = tqdm(train_batches, desc=f"Epoch {epoch}", leave=True)

    for batch_idx, batch in enumerate(train_pbar):
        # ========== 步骤1: 用当前策略采样(On-Policy) ==========
        # 关键:使用当前策略 student_train 生成数据
        with torch.no_grad():
            full_seq, prompt_len, student_old_logp = generate_and_compute_logp(student_train, batch)
            teacher_logp_full = get_teacher_logprob(teacher_model, full_seq, device_teacher)
            # 提取teacher的生成部分
            teacher_logp = teacher_logp_full[:, prompt_len:]

        # ========== 步骤2: 计算当前策略的logp(有梯度) ==========
        student_train.train()
        student_out = student_train(full_seq)
        student_new_logsoft = F.log_softmax(student_out.logits, dim=-1)

        # 只取生成部分的logp
        generated_ids = full_seq[:, prompt_len:]
        student_new_logp = torch.gather(
            student_new_logsoft[:, prompt_len:, :],
            -1,
            generated_ids.unsqueeze(-1)
        ).squeeze(-1)

        # ========== 步骤3: 计算OPD loss ==========
        # 直接最小化 student 和 teacher 的 KL 散度
        loss = compute_opd_loss(student_new_logp, teacher_logp)
        loss = loss / GRAD_ACCUM_STEPS

        # 反向传播更新当前策略
        accelerator.backward(loss)

        # 梯度累积
        if (batch_idx + 1) % GRAD_ACCUM_STEPS == 0:
            if accelerator.sync_gradients:
                torch.nn.utils.clip_grad_norm_(student_train.parameters(), GRAD_CLIP_NORM)
            optimizer.step()
            optimizer.zero_grad()
            global_step += 1

            # 计算监控指标
            with torch.no_grad():
                # 计算KL散度(用于监控)
                min_len = min(student_new_logp.size(1), teacher_logp.size(1))
                s_logp = student_new_logp[:, :min_len]
                t_logp = teacher_logp[:, :min_len]
                kl_value = F.kl_div(
                    F.log_softmax(s_logp, dim=-1),
                    F.softmax(t_logp, dim=-1),
                    reduction='batchmean'
                ).item()

                actual_loss = loss.item() * GRAD_ACCUM_STEPS

            train_pbar.set_postfix({
                "step": global_step,
                "loss": f"{actual_loss:.4f}",
                "kl": f"{kl_value:.4f}"
            })
            write_log(f"【TRAIN】epoch:{epoch}, step:{global_step}, loss={actual_loss:.6f}, kl={kl_value:.6f}")

            # 保存checkpoint
            if global_step % SAVE_STEP_INTERVAL == 0:
                save_path = f"{SAVE_LORA_DIR}/epoch_{epoch}_step_{global_step}"
                unwrapped_model = accelerator.unwrap_model(student_train)
                unwrapped_model.save_pretrained(save_path)
                print(f"\n✅ 保存 checkpoint: {save_path}")
                write_log(f"【SAVE】epoch:{epoch}, step:{global_step}, saved checkpoint")

    # 补齐最后剩余梯度
    if (batch_idx + 1) % GRAD_ACCUM_STEPS != 0:
        if accelerator.sync_gradients:
            torch.nn.utils.clip_grad_norm_(student_train.parameters(), GRAD_CLIP_NORM)
        optimizer.step()
        optimizer.zero_grad()

    # Epoch结束验证
    print(f"\n===== Epoch {epoch} Validation =====")
    val_kl, val_ce = run_validation(student_train, epoch, global_step)
    print(f"KL散度 = {val_kl:.4f} | CE损失 = {val_ce:.4f}")

    if val_kl < best_val_kl:
        best_val_kl = val_kl
        unwrapped_model = accelerator.unwrap_model(student_train)
        unwrapped_model.save_pretrained(BEST_LORA_PATH)
        msg = f"🎉 刷新最优模型,best_KL={best_val_kl:.6f}"
        print(msg)
        write_log(f"【BEST】epoch:{epoch}, step:{global_step}, best_KL={best_val_kl:.6f}")

# 保存最终模型
unwrapped_model = accelerator.unwrap_model(student_train)
unwrapped_model.save_pretrained(f"{SAVE_LORA_DIR}/final_opd_lora")
write_log("========== Training Complete ==========")

print(f"\n训练结束。最优验证KL散度:{best_val_kl:.4f}")
print(f"最终LoRA:{SAVE_LORA_DIR}/final_opd_lora")
print(f"最优LoRA:{BEST_LORA_PATH}")
write_log(f"Final best val KL: {best_val_kl:.6f}")

3、代码解读

(1)模型加载

- 加载student模型,3b模型,用于训练的。

- 加载teacher模型,7b模型,用于蒸馏的。

print("加载 Student Qwen2.5-3B -> GPU0")
student_train = AutoModelForCausalLM.from_pretrained(
    STUDENT_BASE_PATH,
    quantization_config=bnb_config,
    device_map=device_train,
    torch_dtype=torch.bfloat16,
    trust_remote_code=True
)
student_train = prepare_model_for_kbit_training(student_train)
student_train.config.use_cache = False
student_train = get_peft_model(student_train, lora_config)
student_train.print_trainable_parameters()

print("加载 Teacher 7B医疗模型 -> GPU1")
teacher_model = AutoModelForCausalLM.from_pretrained(
    TEACHER_MED_PATH,
    quantization_config=bnb_config,
    device_map=device_teacher,
    torch_dtype=torch.bfloat16,
    trust_remote_code=True
).eval()
teacher_model.config.use_cache = False
for p in teacher_model.parameters():
    p.requires_grad = False

(2)批量推理生成

- 通过model.generate批量生成完整文本序列

- 通过model()生成整条序列的logits

@torch.no_grad()
def generate_and_compute_logp(model, prompt_batch):
    """
    用当前策略生成序列并计算logp(完全无梯度)
    返回:完整序列、prompt长度、生成部分的logp
    """
    model.eval()
    inputs = tokenizer(prompt_batch, return_tensors="pt", padding=True).to(device_train)
    prompt_len = inputs["input_ids"].size(1)

    gen_cfg = GenerationConfig(
        max_new_tokens=MAX_NEW_TOKENS,
        temperature=TEMPERATURE,
        do_sample=True,
        pad_token_id=tokenizer.eos_token_id,
        eos_token_id=tokenizer.eos_token_id,
        num_beams=1,
        use_cache=True
    )
    full_seq = model.generate(**inputs, generation_config=gen_cfg)

    # 提取生成的部分
    generated_ids = full_seq[:, prompt_len:]

    # 计算完整序列的logp
    out = model(full_seq)
    log_soft = F.log_softmax(out.logits, dim=-1)

    # 只取生成部分的logp
    old_logp = torch.gather(
        log_soft[:, prompt_len:, :],
        dim=-1,
        index=generated_ids.unsqueeze(-1)
    ).squeeze(-1)

    return full_seq, prompt_len, old_logp

(3)Teacher模型计算token的对数概率

使用 Teacher 模型,计算整条序列每个位置真实 token 对应的对数概率,传回学生所在设备,用于 OPD 优势计算。

@torch.no_grad()
def get_teacher_logprob(model, seq_ids, target_device):
    """计算teacher模型的log概率"""
    seq_ids = seq_ids.to(target_device)
    out = model(seq_ids)
    log_soft = F.log_softmax(out.logits, dim=-1)
    teacher_logp = torch.gather(log_soft, dim=-1, index=seq_ids.unsqueeze(-1)).squeeze(-1)
    return teacher_logp.to(device_train)

(4)计算loss

def compute_opd_loss(student_logp, teacher_logp):
    """
    计算OPD loss:直接最小化student和teacher的KL散度
    - student_logp: 当前策略在生成部分的logp(有梯度)
    - teacher_logp: teacher模型在生成部分的logp(无梯度)
    """
    # 维度对齐
    min_len = min(student_logp.size(1), teacher_logp.size(1))
    s_logp = student_logp[:, :min_len]
    t_logp = teacher_logp[:, :min_len]

    # 计算KL散度:KL(student || teacher)
    # 注意:这里student是当前策略,teacher是目标分布
    kl_loss = F.kl_div(
        F.log_softmax(s_logp, dim=-1),
        F.softmax(t_logp, dim=-1),
        reduction='batchmean'
    )

    return kl_loss

4、运行日志

[2026-08-03 21:26:13] ========== Training Start ==========
[2026-08-03 21:26:13] ========== Epoch 0 Start ==========
[2026-08-03 21:28:31] 【TRAIN】epoch:0, step:1, loss=2.193159, kl=2.193159
[2026-08-03 21:30:50] 【TRAIN】epoch:0, step:2, loss=3.828450, kl=3.828450
[2026-08-03 21:33:07] 【TRAIN】epoch:0, step:3, loss=1.849946, kl=1.849946
[2026-08-03 21:35:24] 【TRAIN】epoch:0, step:4, loss=2.498603, kl=2.498603
[2026-08-03 21:37:40] 【TRAIN】epoch:0, step:5, loss=2.319613, kl=2.319613
[2026-08-03 21:40:00] 【TRAIN】epoch:0, step:6, loss=3.172891, kl=3.172891
[2026-08-03 21:42:18] 【TRAIN】epoch:0, step:7, loss=2.023333, kl=2.023333
[2026-08-03 21:44:33] 【TRAIN】epoch:0, step:8, loss=1.193600, kl=1.193600
[2026-08-03 21:46:50] 【TRAIN】epoch:0, step:9, loss=1.551316, kl=1.551316
[2026-08-03 21:49:08] 【TRAIN】epoch:0, step:10, loss=1.089417, kl=1.089417
[2026-08-03 21:51:26] 【TRAIN】epoch:0, step:11, loss=1.005224, kl=1.005224
[2026-08-03 21:53:41] 【TRAIN】epoch:0, step:12, loss=1.655025, kl=1.655025
[2026-08-03 21:55:58] 【TRAIN】epoch:0, step:13, loss=1.419596, kl=1.419596
[2026-08-03 21:58:14] 【TRAIN】epoch:0, step:14, loss=3.394303, kl=3.394303
[2026-08-03 22:00:28] 【TRAIN】epoch:0, step:15, loss=1.437412, kl=1.437412
[2026-08-03 22:02:46] 【TRAIN】epoch:0, step:16, loss=2.398520, kl=2.398520
[2026-08-03 22:05:00] 【TRAIN】epoch:0, step:17, loss=5.142252, kl=5.142252
[2026-08-03 22:14:41] 【VAL】epoch:0, step:17, KL=-2.480229, CE=15.582500
[2026-08-03 22:14:41] 【BEST】epoch:0, step:17, best_KL=-2.480229
[2026-08-03 22:14:41] ========== Epoch 1 Start ==========
[2026-08-03 22:16:56] 【TRAIN】epoch:1, step:18, loss=1.228516, kl=1.228516
[2026-08-03 22:19:10] 【TRAIN】epoch:1, step:19, loss=1.636437, kl=1.636437
[2026-08-03 22:21:24] 【TRAIN】epoch:1, step:20, loss=0.710897, kl=0.710897
[2026-08-03 22:23:40] 【TRAIN】epoch:1, step:21, loss=1.929194, kl=1.929194
[2026-08-03 22:25:55] 【TRAIN】epoch:1, step:22, loss=2.360214, kl=2.360214
[2026-08-03 22:28:12] 【TRAIN】epoch:1, step:23, loss=1.857954, kl=1.857954
[2026-08-03 22:30:29] 【TRAIN】epoch:1, step:24, loss=1.244854, kl=1.244854
[2026-08-03 22:32:47] 【TRAIN】epoch:1, step:25, loss=2.152761, kl=2.152761
[2026-08-03 22:35:03] 【TRAIN】epoch:1, step:26, loss=1.143138, kl=1.143138
[2026-08-03 22:37:19] 【TRAIN】epoch:1, step:27, loss=3.202857, kl=3.202857
[2026-08-03 22:39:34] 【TRAIN】epoch:1, step:28, loss=1.311961, kl=1.311961
[2026-08-03 22:41:52] 【TRAIN】epoch:1, step:29, loss=3.203570, kl=3.203570
[2026-08-03 22:44:07] 【TRAIN】epoch:1, step:30, loss=2.545751, kl=2.545751
[2026-08-03 22:46:22] 【TRAIN】epoch:1, step:31, loss=1.632395, kl=1.632395
[2026-08-03 22:48:38] 【TRAIN】epoch:1, step:32, loss=1.711337, kl=1.711337
[2026-08-03 22:50:51] 【TRAIN】epoch:1, step:33, loss=1.611026, kl=1.611026
[2026-08-03 22:53:04] 【TRAIN】epoch:1, step:34, loss=2.426403, kl=2.426403
[2026-08-03 23:03:01] 【VAL】epoch:1, step:34, KL=-3.372311, CE=15.415000
[2026-08-03 23:03:01] 【BEST】epoch:1, step:34, best_KL=-3.372311
[2026-08-03 23:03:01] ========== Epoch 2 Start ==========
[2026-08-03 23:05:18] 【TRAIN】epoch:2, step:35, loss=2.039419, kl=2.039419
[2026-08-03 23:07:31] 【TRAIN】epoch:2, step:36, loss=2.615076, kl=2.615076
[2026-08-03 23:09:45] 【TRAIN】epoch:2, step:37, loss=1.750095, kl=1.750095
[2026-08-03 23:11:57] 【TRAIN】epoch:2, step:38, loss=1.769925, kl=1.769925
[2026-08-03 23:14:09] 【TRAIN】epoch:2, step:39, loss=0.852849, kl=0.852849
[2026-08-03 23:16:22] 【TRAIN】epoch:2, step:40, loss=1.551821, kl=1.551821
[2026-08-03 23:18:38] 【TRAIN】epoch:2, step:41, loss=1.044877, kl=1.044877
[2026-08-03 23:20:59] 【TRAIN】epoch:2, step:42, loss=1.569338, kl=1.569338
[2026-08-03 23:23:19] 【TRAIN】epoch:2, step:43, loss=0.930750, kl=0.930750
[2026-08-03 23:25:40] 【TRAIN】epoch:2, step:44, loss=2.612155, kl=2.612155
[2026-08-03 23:27:58] 【TRAIN】epoch:2, step:45, loss=1.342157, kl=1.342157
[2026-08-03 23:30:21] 【TRAIN】epoch:2, step:46, loss=1.394093, kl=1.394093
[2026-08-03 23:32:51] 【TRAIN】epoch:2, step:47, loss=1.500656, kl=1.500656
[2026-08-03 23:35:19] 【TRAIN】epoch:2, step:48, loss=1.236770, kl=1.236770
[2026-08-03 23:37:46] 【TRAIN】epoch:2, step:49, loss=0.614521, kl=0.614521
[2026-08-03 23:40:13] 【TRAIN】epoch:2, step:50, loss=2.114708, kl=2.114708
[2026-08-03 23:42:35] 【TRAIN】epoch:2, step:51, loss=1.013776, kl=1.013776
[2026-08-03 23:53:06] 【VAL】epoch:2, step:51, KL=-2.147025, CE=15.612500
[2026-08-03 23:53:06] ========== Epoch 3 Start ==========
[2026-08-03 23:55:32] 【TRAIN】epoch:3, step:52, loss=2.307380, kl=2.307380
[2026-08-03 23:58:00] 【TRAIN】epoch:3, step:53, loss=1.230930, kl=1.230930
[2026-08-04 00:00:29] 【TRAIN】epoch:3, step:54, loss=0.509488, kl=0.509488
[2026-08-04 00:02:54] 【TRAIN】epoch:3, step:55, loss=0.886873, kl=0.886873
[2026-08-04 00:05:23] 【TRAIN】epoch:3, step:56, loss=1.595298, kl=1.595298
[2026-08-04 00:07:49] 【TRAIN】epoch:3, step:57, loss=1.105819, kl=1.105819
[2026-08-04 00:10:17] 【TRAIN】epoch:3, step:58, loss=1.885654, kl=1.885654
[2026-08-04 00:12:45] 【TRAIN】epoch:3, step:59, loss=1.136614, kl=1.136614
[2026-08-04 00:15:13] 【TRAIN】epoch:3, step:60, loss=1.031398, kl=1.031398
[2026-08-04 00:17:40] 【TRAIN】epoch:3, step:61, loss=1.285468, kl=1.285468
[2026-08-04 00:20:08] 【TRAIN】epoch:3, step:62, loss=0.798619, kl=0.798619
[2026-08-04 00:22:36] 【TRAIN】epoch:3, step:63, loss=0.814253, kl=0.814253
[2026-08-04 00:25:04] 【TRAIN】epoch:3, step:64, loss=0.842485, kl=0.842485
[2026-08-04 00:27:34] 【TRAIN】epoch:3, step:65, loss=1.589093, kl=1.589093
[2026-08-04 00:30:04] 【TRAIN】epoch:3, step:66, loss=1.203898, kl=1.203898
[2026-08-04 00:32:33] 【TRAIN】epoch:3, step:67, loss=1.081560, kl=1.081560
[2026-08-04 00:35:00] 【TRAIN】epoch:3, step:68, loss=0.909622, kl=0.909622
[2026-08-04 00:45:37] 【VAL】epoch:3, step:68, KL=-2.257676, CE=15.540000
[2026-08-04 00:45:37] ========== Epoch 4 Start ==========
[2026-08-04 00:48:02] 【TRAIN】epoch:4, step:69, loss=1.282916, kl=1.282916
[2026-08-04 00:50:30] 【TRAIN】epoch:4, step:70, loss=1.076593, kl=1.076593
[2026-08-04 00:52:58] 【TRAIN】epoch:4, step:71, loss=1.011389, kl=1.011389
[2026-08-04 00:55:27] 【TRAIN】epoch:4, step:72, loss=0.893716, kl=0.893716
[2026-08-04 00:57:58] 【TRAIN】epoch:4, step:73, loss=1.543352, kl=1.543352
[2026-08-04 01:00:27] 【TRAIN】epoch:4, step:74, loss=1.816490, kl=1.816490
[2026-08-04 01:02:57] 【TRAIN】epoch:4, step:75, loss=0.846386, kl=0.846386
[2026-08-04 01:05:29] 【TRAIN】epoch:4, step:76, loss=0.566961, kl=0.566961
[2026-08-04 01:07:57] 【TRAIN】epoch:4, step:77, loss=0.988257, kl=0.988257
[2026-08-04 01:10:26] 【TRAIN】epoch:4, step:78, loss=1.026584, kl=1.026584
[2026-08-04 01:12:56] 【TRAIN】epoch:4, step:79, loss=1.071776, kl=1.071776
[2026-08-04 01:15:25] 【TRAIN】epoch:4, step:80, loss=0.806434, kl=0.806434
[2026-08-04 01:17:54] 【TRAIN】epoch:4, step:81, loss=1.320720, kl=1.320720
[2026-08-04 01:20:26] 【TRAIN】epoch:4, step:82, loss=2.582267, kl=2.582267
[2026-08-04 01:22:56] 【TRAIN】epoch:4, step:83, loss=1.295375, kl=1.295375
[2026-08-04 01:25:25] 【TRAIN】epoch:4, step:84, loss=0.698409, kl=0.698409
[2026-08-04 01:27:49] 【TRAIN】epoch:4, step:85, loss=1.063825, kl=1.063825
[2026-08-04 01:38:31] 【VAL】epoch:4, step:85, KL=-2.317143, CE=15.745000
[2026-08-04 01:38:31] ========== Epoch 5 Start ==========
[2026-08-04 01:41:02] 【TRAIN】epoch:5, step:86, loss=1.258141, kl=1.258141
[2026-08-04 01:43:33] 【TRAIN】epoch:5, step:87, loss=1.013202, kl=1.013202
[2026-08-04 01:46:03] 【TRAIN】epoch:5, step:88, loss=1.027144, kl=1.027144
[2026-08-04 01:48:36] 【TRAIN】epoch:5, step:89, loss=0.705286, kl=0.705286
[2026-08-04 01:51:05] 【TRAIN】epoch:5, step:90, loss=1.039880, kl=1.039880
[2026-08-04 01:53:36] 【TRAIN】epoch:5, step:91, loss=1.964346, kl=1.964346
[2026-08-04 01:56:10] 【TRAIN】epoch:5, step:92, loss=0.857443, kl=0.857443
[2026-08-04 01:58:42] 【TRAIN】epoch:5, step:93, loss=2.905995, kl=2.905995
[2026-08-04 02:01:12] 【TRAIN】epoch:5, step:94, loss=1.182123, kl=1.182123
[2026-08-04 02:03:44] 【TRAIN】epoch:5, step:95, loss=0.973629, kl=0.973629
[2026-08-04 02:06:13] 【TRAIN】epoch:5, step:96, loss=1.984335, kl=1.984335
[2026-08-04 02:08:43] 【TRAIN】epoch:5, step:97, loss=1.586762, kl=1.586762
[2026-08-04 02:11:09] 【TRAIN】epoch:5, step:98, loss=0.830937, kl=0.830937
[2026-08-04 02:13:36] 【TRAIN】epoch:5, step:99, loss=0.728865, kl=0.728865
[2026-08-04 02:16:04] 【TRAIN】epoch:5, step:100, loss=0.890922, kl=0.890922
[2026-08-04 02:18:36] 【TRAIN】epoch:5, step:101, loss=1.150494, kl=1.150494
[2026-08-04 02:21:03] 【TRAIN】epoch:5, step:102, loss=0.909000, kl=0.909000
[2026-08-04 02:31:56] 【VAL】epoch:5, step:102, KL=-2.443864, CE=15.302500
[2026-08-04 02:31:56] ========== Training Complete ==========
[2026-08-04 02:31:56] Final best val KL: -3.372311

评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值