CodeFormer预训练模型加载及API部署问题求助
模块导入错误解决与API部署方案
一、模块导入错误修复步骤
- 修正路径配置
将sys.path.append('*')中的*替换为CodeFormer项目根目录的绝对路径,比如sys.path.append('/home/your-name/CodeFormer'),确保Python能识别项目内的basicsr模块。 - 匹配依赖版本
执行CodeFormer根目录下的依赖安装命令,避免第三方basicsr包版本冲突:
确保PyTorch版本≥1.7.0,torchvision≥0.8.1。pip install -r requirements.txt - 修正模型初始化代码
原代码中CodeFormerModel(__name__)的实例化方式错误,官方模型需要配合配置文件加载,且预训练权重需读取params_ema键:import sys sys.path.append('/path/to/CodeFormer') import torch from basicsr.models import create_model from basicsr.utils.options import parse_options # 加载官方配置文件 opt, _ = parse_options('/path/to/CodeFormer', opt_path='configs/codeformer/codeformer.yaml', is_train=False) model = create_model(opt) # 加载预训练权重 model_path = '/path/to/pretrained/model.pth' model.load_state_dict(torch.load(model_path, map_location=torch.device('cpu'))['params_ema']) model.eval()
二、API部署实现(基于FastAPI)
- 安装依赖
pip install fastapi uvicorn pillow - 编写API服务代码(
main.py)import sys sys.path.append('/path/to/CodeFormer') import torch from fastapi import FastAPI, File, UploadFile from PIL import Image from basicsr.models import create_model from basicsr.utils.options import parse_options from basicsr.utils import img2tensor, tensor2img import base64 from io import BytesIO app = FastAPI() # 全局初始化模型(仅加载一次) opt, _ = parse_options('/path/to/CodeFormer', opt_path='configs/codeformer/codeformer.yaml', is_train=False) model = create_model(opt) model.load_state_dict(torch.load('/path/to/pretrained/model.pth', map_location=torch.device('cpu'))['params_ema']) model.eval() @app.post("/enhance-image") async def enhance_image(file: UploadFile = File(...)): # 读取并预处理输入图片 img = Image.open(file.file).convert('RGB') img_tensor = img2tensor(img, bgr2rgb=True, float32=True).unsqueeze(0) # 模型推理 with torch.no_grad(): output = model(img_tensor, w=0.5)['output'] # w为修复强度,取值0-1 # 转换结果为base64返回 output_img = tensor2img(output.squeeze(0), rgb2bgr=True) output_pil = Image.fromarray(output_img) buffer = BytesIO() output_pil.save(buffer, format='JPEG') img_str = base64.b64encode(buffer.getvalue()).decode() return {"enhanced_image_base64": img_str} - 启动服务
启动后可访问uvicorn main:app --host 0.0.0.0 --port 8000http://localhost:8000/docs测试API接口。
内容的提问来源于stack exchange,提问作者Cyan
相关产品推荐
相关产品推荐

