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

MLFlow技术问询:如何加载模型权重并继续训练(含微调)

从MLFlow加载模型并以不同学习率继续训练

核心流程

从MLFlow加载预训练模型 → 配置新学习率的优化器 → 继续执行训练


1. 加载MLFlow存储的模型

根据你使用的框架,选择对应的MLFlow加载方法:

TensorFlow/Keras模型

import mlflow.tensorflow
import tensorflow as tf

# 替换为你的模型URI,支持模型注册表或特定运行路径
model_uri = "models:/your_model_name/1"  # 示例:版本1的模型
loaded_model = mlflow.tensorflow.load_model(model_uri)

PyTorch模型

import mlflow.pytorch
import torch

model_uri = "models:/your_model_name/1"
loaded_model = mlflow.pytorch.load_model(model_uri)

2. 配置带新学习率的优化器

继续训练时需重新定义优化器,指定目标学习率(微调通常用预训练时10-100倍小的学习率):

TensorFlow/Keras示例

new_lr = 1e-5  # 自定义新学习率
optimizer = tf.keras.optimizers.Adam(learning_rate=new_lr)

# 沿用原模型的损失函数和评估指标重新编译
loaded_model.compile(
    optimizer=optimizer,
    loss=loaded_model.loss,
    metrics=loaded_model.metrics_names
)

PyTorch示例

new_lr = 1e-5
optimizer = torch.optim.Adam(loaded_model.parameters(), lr=new_lr)

# (可选)若需延续原优化器状态,需在训练时手动log优化器state_dict到MLFlow,加载时恢复
# optimizer_state = torch.load(mlflow.artifacts.download_artifacts("path/to/optimizer_state.pth"))
# optimizer.load_state_dict(optimizer_state)

3. 启动继续训练

TensorFlow/Keras示例

# 假设已准备好训练/验证数据集
history = loaded_model.fit(
    train_dataset,
    validation_data=val_dataset,
    epochs=20,  # 新增训练轮次
    initial_epoch=loaded_model.history.epoch[-1] if hasattr(loaded_model, 'history') else 0
)

PyTorch示例

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
loaded_model.to(device)
criterion = torch.nn.CrossEntropyLoss()  # 替换为你的损失函数

epochs = 10
for epoch in range(epochs):
    loaded_model.train()
    running_loss = 0.0
    for inputs, labels in train_loader:
        inputs, labels = inputs.to(device), labels.to(device)
        
        optimizer.zero_grad()
        outputs = loaded_model(inputs)
        loss = criterion(outputs, labels)
        loss.backward()
        optimizer.step()
        
        running_loss += loss.item() * inputs.size(0)
    
    print(f"Epoch {epoch+1}/{epochs}, Train Loss: {running_loss/len(train_loader.dataset):.4f}")

    # 可选验证步骤
    loaded_model.eval()
    val_loss = 0.0
    with torch.no_grad():
        for inputs, labels in val_loader:
            inputs, labels = inputs.to(device), labels.to(device)
            outputs = loaded_model(inputs)
            val_loss += criterion(outputs, labels).item() * inputs.size(0)
    print(f"Validation Loss: {val_loss/len(val_loader.dataset):.4f}")

关键注意事项

  • 模型URI格式:支持models:/<模型名>/<版本号>(模型注册表)或runs:/<运行ID>/<artifact路径>(单运行的模型文件)
  • 版本兼容性:确保加载模型的框架版本与训练时一致,避免权重加载失败
  • 学习率策略:微调大模型时,可尝试分层设置学习率(仅解冻顶层参数用大学习率,底层用极小学习率)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 16:55:27