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

如何基于Azure Durable Functions实现MNIST分类的分步训练

解决Azure Durable Functions + MNIST训练的核心问题

一、修正Activity函数间数据传递问题

Durable Functions依赖JSON序列化传递数据,直接传递PyTorch模型/优化器对象会失败,必须传递可序列化的状态数据。核心修正点如下:

  • 传递模型的state_dict()(字典格式,可直接序列化)而非模型对象
  • 传递优化器的state_dict()而非优化器对象
  • 所有Activity的输入输出仅用基础数据类型(字典、数值、列表等)

修正后的训练Activity示例

import torch
from torch.utils.data import DataLoader
from torchvision.datasets import MNIST
from torchvision.transforms import ToTensor
from your_module import Net  # 复用现有Net类

def train_activity(input_data):
    # 解析传入的序列化参数
    model_state = input_data.get("model_state")
    optimizer_state = input_data.get("optimizer_state")
    num_epochs = input_data["num_epochs"]
    batch_size = input_data.get("batch_size", 64)

    # 初始化模型与优化器,加载历史状态
    model = Net()
    if model_state:
        model.load_state_dict(model_state)
    optimizer = torch.optim.Adam(model.parameters())
    if optimizer_state:
        optimizer.load_state_dict(optimizer_state)

    # 加载MNIST数据集
    train_dataset = MNIST(root="./data", train=True, download=True, transform=ToTensor())
    train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)

    # 执行训练逻辑
    model.train()
    total_loss = 0.0
    for epoch in range(num_epochs):
        epoch_loss = 0.0
        for data, target in train_loader:
            optimizer.zero_grad()
            output = model(data)
            loss = torch.nn.functional.cross_entropy(output, target)
            loss.backward()
            optimizer.step()
            epoch_loss += loss.item()
        total_loss = epoch_loss / len(train_loader)
        print(f"完成第{epoch+1}轮训练,损失值:{total_loss:.4f}")

    # 返回序列化的训练结果
    return {
        "model_state": model.state_dict(),
        "optimizer_state": optimizer.state_dict(),
        "final_loss": total_loss
    }

二、实现Epoch分步执行

在Orchestrator函数中拆分训练任务,分两次调用训练Activity,传递不同的epoch数,并将前一次训练的状态传入下一次任务,实现续训。

修正后的Orchestrator示例

import azure.durable_functions as df

def orchestrator_function(context: df.DurableOrchestrationContext):
    # 第一步:执行3轮训练
    first_train_params = {
        "model_state": None,
        "optimizer_state": None,
        "num_epochs": 3
    }
    first_train_result = yield context.call_activity("TrainActivity", first_train_params)

    # 第二步:基于前3轮的状态,继续执行2轮训练
    second_train_params = {
        "model_state": first_train_result["model_state"],
        "optimizer_state": first_train_result["optimizer_state"],
        "num_epochs": 2
    }
    second_train_result = yield context.call_activity("TrainActivity", second_train_params)

    # 可选:调用评估Activity验证最终模型
    eval_result = yield context.call_activity("EvaluateActivity", {
        "model_state": second_train_result["model_state"]
    })

    # 返回最终汇总结果
    return {
        "总训练轮数": 3 + 2,
        "最终训练损失": second_train_result["final_loss"],
        "测试集准确率": eval_result.get("accuracy")
    }

main = df.Orchestrator.create(orchestrator_function)

额外注意事项

  • 确保Azure Functions环境安装PyTorch、torchvision依赖,需在requirements.txt中添加对应包
  • Durable Functions更适合小规模分步训练场景,大规模训练建议使用专用ML训练平台
  • 所有Activity函数需保持无状态,仅依赖传入的参数执行逻辑

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.10 21:55:06