如何基于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
相关产品推荐
相关产品推荐

