加载MLflow存储的PyTorch Lightning模型后predict方法报错,求解决方案
问题
我训练了一个PyTorch Lightning模型,使用MLflow记录并成功加载后,调用predict方法时出现错误。
我的代码:
classifier_model = TextClassifier(backbone="prajjwal1/bert-tiny", num_classes=datamodule.num_classes, metrics=torchmetrics.F1Score(datamodule.num_classes)) trainer = flash.Trainer(max_epochs=3, gpus=torch.cuda.device_count()) MODEL_ARTIFACT_PATH = 'MODEL' REGISTERED_MODEL_NAME = 'MODEL2' with mlflow.start_run(experiment_id=experiment.experiment_id, run_name="MYRUN01") as dl_model_tracking_run: trainer.finetune(classifier_model, datamodule=datamodule, strategy="freeze") trainer.test(dataloaders=datamodule) .... mlflow.pytorch.log_model(pytorch_model=classifier_model, artifact_path=MODEL_ARTIFACT_PATH, registered_model_name=REGISTERED_MODEL_NAME) run_id = dl_model_tracking_run.info.run_id print("run_id: {}; lifecycle_stage: {}".format(run_id, mlflow.get_run(run_id).info.lifecycle_stage)) logged_model = f'runs:/{run_id}/{MODEL_ARTIFACT_PATH}' model = mlflow.pytorch.load_model(logged_model) model.trainer.state.stage='test' model.predict({'What a news!'})
错误信息:
def predict(self, *args, **kwargs): raise AttributeError("`flash.Task.predict` has been removed. Use `flash.Trainer.predict` instead.") AttributeError: `flash.Task.predict` has been removed. Use `flash.Trainer.predict` instead.
我查阅了MLflow官方文档,认为代码并无问题,请问该如何解决此问题?
解决方法
错误提示已明确指出:flash.Task.predict 方法已被移除,必须改用 flash.Trainer.predict 方法。具体修改如下:
- 加载模型后,创建(或复用)
flash.Trainer实例 - 调用
trainer.predict()方法,传入加载后的模型和预测数据
修改后的代码片段:
logged_model = f'runs:/{run_id}/{MODEL_ARTIFACT_PATH}' model = mlflow.pytorch.load_model(logged_model) # 创建Trainer实例 trainer = flash.Trainer(gpus=torch.cuda.device_count()) # 使用Trainer的predict方法执行预测 predictions = trainer.predict(model, data={'What a news!'})
额外注意事项:
- 确保传入
trainer.predict()的数据格式与模型训练时的输入格式一致,避免因格式不匹配导致新的错误 - 由于Flash框架API更新,所有Flash Task类的预测操作都必须通过Trainer来执行,不能再直接调用Task自身的predict方法
内容的提问来源于stack exchange,提问作者Patiasa
相关产品推荐
相关产品推荐

