如何在Beam中结合自定义PyTorch模型使用RunInference?
在GCP Dataflow(apache-beam-with-gcp=2.44.0)中适配自定义PyTorch模型到RunInference的问题与解决
问题背景
在GCP上使用apache-beam-with-gcp=2.44.0版本的Dataflow,基于自定义PyTorch类构建机器学习模型,本地加载模型的代码如下:
input_model = path / "model.ckpt" checkpoint = torch.load(input_model, map_location=torch.device("cpu")) model_name = checkpoint["parameters"]["model_name"] n_classes = len(checkpoint["parameters"]["class_names"]) backbone_fn = model_architecture.get_backbone(model_name) backbone = backbone_fn(num_classes=n_classes) model = CustomModel.load_from_checkpoint( checkpoint_path=input_model, backbone=backbone )
尝试按照官方文档使用RunInference组件执行推理,代码如下:
with pipeline as p: ( p | "ReadInputData" >> beam.Create(value_to_predict) | "RunInferenceTorch" >> RunInference(torch_model_handler) | beam.Map(print) )
但使用PytorchModelHandlerTensor适配自定义模型时失败,最初的Handler配置代码:
torch_model_handler = PytorchModelHandlerTensor( state_dict_path=None, model_class=CustomModel.load_from_checkpoint, model_params={"backbone": backbone, "checkpoint_path": input_model, })
问题原因
查看apache_beam/ml/inference/pytorch_inference.py中的模型加载逻辑,其核心流程为:
model.load_state_dict(state_dict) model.to(device) model.eval()
该逻辑与自定义模型通过load_from_checkpoint的加载方式不兼容——Handler预期先初始化模型实例,再加载state_dict,而直接传入CustomModel.load_from_checkpoint作为model_class不符合这个标准流程。
解决方案
先加载模型并将其保存为标准的PyTorch state_dict格式,再使用PytorchModelHandlerTensor正常加载:
# 先加载模型并保存state_dict torch.save(model.state_dict(), './model_torch.pt') # 配置PytorchModelHandlerTensor torch_model_handler = PytorchModelHandlerTensor( state_dict_path='./model_torch.pt', model_class=CustomModel, model_params={"backbone": backbone, "criterion": criterion, "class_names": class_names, "expected_input_size": expected_input_size })
内容的提问来源于stack exchange,提问作者Dr. Fabien Tarrade
相关产品推荐
相关产品推荐

