端到端图像复原系统实战:从 API 服务到 Docker 部署的完整工程方案

一、引言

前九篇文章我们深入学习了 Real-ESRGAN、DehazeFormer 和 Restormer 的原理与实战。本文将这些模型整合为一个生产级图像复原微服务,覆盖 API 设计、模型管理、异步处理、Docker 部署等工程实践。

二、系统架构总览

                    ┌──────────────┐
                    │   Nginx LB   │  (反向代理 / HTTPS / 限流)
                    └──────┬───────┘
                           │
              ┌────────────┼────────────┐
              ▼            ▼            ▼
        ┌──────────┐ ┌──────────┐ ┌──────────┐
        │ FastAPI  │ │ FastAPI  │ │ FastAPI  │
        │ Worker 1 │ │ Worker 2 │ │ Worker N │
        └────┬─────┘ └────┬─────┘ └────┬─────┘
             │            │            │
             └────────────┼────────────┘
                          │
                    ┌─────▼──────┐
                    │   Redis    │  (任务队列 + 结果缓存)
                    └─────┬──────┘
                          │
                    ┌─────▼──────┐
                    │   Celery   │  (异步任务处理)
                    │   Workers  │
                    └─────┬──────┘
                          │
              ┌───────────┼───────────┐
              ▼           ▼           ▼
        ┌──────────┐ ┌──────────┐ ┌──────────┐
        │ Real     │ │ Dehaze   │ │ Restormer│
        │ ESRGAN   │ │ Former   │ │          │
        │ (GPU 0)  │ │ (GPU 1)  │ │ (GPU 2)  │
        └──────────┘ └──────────┘ └──────────┘

三、FastAPI 服务实现

3.1 项目结构

image_restoration_api/
├── app/
│   ├── __init__.py
│   ├── main.py              # FastAPI 入口
│   ├── config.py            # 配置管理
│   ├── models/
│   │   ├── __init__.py
│   │   ├── manager.py       # 模型管理器
│   │   ├── super_res.py     # Real-ESRGAN 封装
│   │   ├── dehaze.py        # DehazeFormer 封装
│   │   └── denoise.py       # Restormer 封装
│   ├── pipeline.py          # 管线编排
│   ├── schemas.py           # Pydantic 请求/响应模型
│   └── utils.py             # 工具函数
├── tests/
│   └── test_api.py
├── Dockerfile
├── docker-compose.yml
├── requirements.txt
└── .env

3.2 请求/响应模型

# app/schemas.py
from pydantic import BaseModel, Field
from typing import Optional, List
from enum import Enum
from datetime import datetime


class TaskType(str, Enum):
    SUPER_RESOLUTION = "super_resolution"
    DENOISE = "denoise"
    DEHAZE = "dehaze"
    DEBLUR = "deblur"
    DERAIN = "derain"
    AUTO = "auto"


class RestoreRequest(BaseModel):
    """图像复原请求"""
    tasks: List[TaskType] = Field(
        default=[TaskType.AUTO],
        description="处理任务列表,AUTO 表示自动检测"
    )
    options: Optional[dict] = Field(
        default=None,
        description="任务参数,如 {'scale': 4, 'noise_level': 25}"
    )


class RestoreResponse(BaseModel):
    """异步任务响应"""
    task_id: str
    status: str = "processing"
    created_at: datetime


class TaskResult(BaseModel):
    """任务结果"""
    task_id: str
    status: str  # completed / failed
    result_url: Optional[str] = None
    processing_time: Optional[float] = None
    error: Optional[str] = None
    degradation_diagnosis: Optional[dict] = None

3.3 模型管理器

# app/models/manager.py
import torch
import threading
import time
from typing import Dict, Any


class ModelManager:
    """模型生命周期管理器"""

    _instance = None
    _lock = threading.Lock()

    def __new__(cls):
        if cls._instance is None:
            with cls._lock:
                if cls._instance is None:
                    cls._instance = super().__new__(cls)
                    cls._instance._initialized = False
        return cls._instance

    def __init__(self):
        if self._initialized:
            return
        self._initialized = True

        self.models: Dict[str, Any] = {}
        self.device = torch.device(
            'cuda' if torch.cuda.is_available() else 'cpu'
        )
        self._model_load_times: Dict[str, float] = {}

    def get_model(self, task: str):
        """获取或懒加载模型"""
        if task not in self.models:
            self.models[task] = self._load_model(task)
        return self.models[task]

    def _load_model(self, task: str):
        """加载指定任务的模型"""
        print(f"加载模型: {task}")
        t_start = time.time()

        if task == 'super_resolution':
            from app.models.super_res import SuperResolutionModel
            model = SuperResolutionModel(device=self.device)

        elif task == 'denoise':
            from app.models.denoise import DenoiseModel
            model = DenoiseModel(device=self.device)

        elif task == 'dehaze':
            from app.models.dehaze import DehazeModel
            model = DehazeModel(device=self.device)

        elif task == 'deblur':
            from app.models.denoise import DeblurModel
            model = DeblurModel(device=self.device)

        elif task == 'derain':
            from app.models.dehaze import DerainModel
            model = DerainModel(device=self.device)

        else:
            raise ValueError(f"未知任务: {task}")

        self._model_load_times[task] = time.time() - t_start
        print(f"模型 {task} 加载完成 ({self._model_load_times[task]:.1f}s)")

        return model

    def warmup(self):
        """预热所有模型"""
        tasks = ['super_resolution', 'denoise', 'dehaze', 'deblur']
        for task in tasks:
            self.get_model(task)

    def clear_cache(self):
        """清理 GPU 缓存"""
        if self.device.type == 'cuda':
            torch.cuda.empty_cache()

3.4 管线编排

# app/pipeline.py
import time
import numpy as np
import cv2
from app.models.manager import ModelManager
from app.utils import ImageAnalyzer


class RestorationPipeline:
    """图像复原管线"""

    TASK_ORDER = ['denoise', 'deblur', 'derain', 'dehaze', 'super_resolution']

    def __init__(self):
        self.model_manager = ModelManager()
        self.analyzer = ImageAnalyzer()

    def process(self, image_bytes: bytes, tasks: list,
                options: dict = None) -> dict:
        """处理图像"""
        t_start = time.time()

        # 1. 解码图像
        img_array = np.frombuffer(image_bytes, np.uint8)
        img = cv2.imdecode(img_array, cv2.IMREAD_COLOR)
        if img is None:
            raise ValueError("无法解码图像")

        original_shape = img.shape
        print(f"输入: {original_shape[1]}×{original_shape[0]}")

        # 2. 退化诊断
        diagnosis = self.analyzer.analyze(img)

        # 3. 自动检测任务
        if 'auto' in [t.value if hasattr(t, 'value') else t for t in tasks] or            'auto' in tasks:
            effective_tasks = self._auto_tasks(diagnosis)
        else:
            effective_tasks = [t.value if hasattr(t, 'value') else t
                              for t in tasks]

        # 4. 按顺序执行
        result = img
        task_details = []

        for task in self.TASK_ORDER:
            if task not in effective_tasks:
                continue

            model = self.model_manager.get_model(task)

            # 执行任务
            task_start = time.time()
            result = model.process(result, options or {})
            task_time = time.time() - task_start

            task_details.append({
                'task': task,
                'time_ms': round(task_time * 1000, 2),
                'output_shape': list(result.shape)
            })

        # 5. 编码输出
        success, encoded = cv2.imencode('.png', result)
        if not success:
            raise RuntimeError("编码输出失败")

        total_time = time.time() - t_start

        return {
            'status': 'completed',
            'output_bytes': encoded.tobytes(),
            'processing_time_ms': round(total_time * 1000, 2),
            'diagnosis': diagnosis,
            'task_details': task_details,
            'input_shape': list(original_shape),
            'output_shape': list(result.shape)
        }

    def _auto_tasks(self, diagnosis):
        """自动确定任务列表"""
        tasks = []
        if diagnosis.get('noise_level', 0) > 0.3:
            tasks.append('denoise')
        if diagnosis.get('blur_score', 200) < 150:
            tasks.append('deblur')
        if diagnosis.get('rain_score', 0) > 0.3:
            tasks.append('derain')
        if diagnosis.get('haze_density', 0) > 0.4:
            tasks.append('dehaze')
        if diagnosis.get('resolution_score', 1) < 0.5:
            tasks.append('super_resolution')
        return tasks if tasks else ['denoise']  # fallback

3.5 FastAPI 路由

# app/main.py
from fastapi import FastAPI, UploadFile, File, HTTPException, Query
from fastapi.responses import Response, JSONResponse
from app.pipeline import RestorationPipeline
from app.schemas import RestoreRequest, TaskResult, TaskType
import uuid
import asyncio
from typing import List, Optional


app = FastAPI(
    title="图像复原服务",
    description="基于 Real-ESRGAN + DehazeFormer + Restormer 的图像复原微服务",
    version="1.0.0"
)

pipeline = RestorationPipeline()

# 异步任务存储(生产环境应使用 Redis)
task_store = {}


@app.on_event("startup")
async def startup():
    """启动时预热模型"""
    # 可选:预加载所有模型
    # pipeline.model_manager.warmup()
    print("服务启动完成")


@app.post("/restore", response_model=TaskResult)
async def restore_image(
    file: UploadFile = File(...),
    tasks: str = Query(default="auto",
                       description="逗号分隔的任务列表"),
    scale: int = Query(default=4, ge=1, le=8),
    noise_level: int = Query(default=25, ge=0, le=100)
):
    """
    图像复原接口
    支持: super_resolution, denoise, dehaze, deblur, derain, auto
    """
    # 验证文件类型
    if not file.content_type.startswith('image/'):
        raise HTTPException(400, "仅支持图像文件")

    # 读取图像
    image_bytes = await file.read()
    if len(image_bytes) > 20 * 1024 * 1024:  # 20MB 限制
        raise HTTPException(400, "图像大小不能超过 20MB")

    # 立即返回 task_id,后台处理
    task_id = str(uuid.uuid4())
    task_store[task_id] = {'status': 'processing'}

    # 异步处理
    asyncio.create_task(
        _process_async(task_id, image_bytes, tasks, scale, noise_level)
    )

    return TaskResult(
        task_id=task_id,
        status="processing"
    )


@app.get("/task/{task_id}", response_model=TaskResult)
async def get_task_result(task_id: str):
    """查询任务结果"""
    if task_id not in task_store:
        raise HTTPException(404, "任务不存在")

    task = task_store[task_id]
    if task['status'] == 'completed':
        return TaskResult(
            task_id=task_id,
            status='completed',
            result_url=f'/download/{task_id}',
            processing_time=task.get('processing_time'),
            degradation_diagnosis=task.get('diagnosis')
        )
    elif task['status'] == 'failed':
        return TaskResult(
            task_id=task_id,
            status='failed',
            error=task.get('error')
        )
    else:
        return TaskResult(task_id=task_id, status='processing')


@app.get("/download/{task_id}")
async def download_result(task_id: str):
    """下载结果图像"""
    if task_id not in task_store:
        raise HTTPException(404, "任务不存在")

    task = task_store[task_id]
    if task['status'] != 'completed':
        raise HTTPException(400, "任务未完成")

    return Response(
        content=task['output_bytes'],
        media_type="image/png",
        headers={
            "Content-Disposition": f"attachment; filename=restored_{task_id}.png"
        }
    )


@app.get("/health")
async def health_check():
    """健康检查"""
    return {
        "status": "healthy",
        "models_loaded": list(pipeline.model_manager.models.keys()),
        "gpu_available": torch.cuda.is_available(),
        "gpu_name": torch.cuda.get_device_name(0) if torch.cuda.is_available() else None
    }


async def _process_async(task_id, image_bytes, tasks_str, scale, noise_level):
    """异步处理图像"""
    try:
        task_list = [t.strip() for t in tasks_str.split(',')]

        result = await asyncio.to_thread(
            pipeline.process,
            image_bytes,
            task_list,
            {'scale': scale, 'noise_level': noise_level}
        )

        task_store[task_id] = result
    except Exception as e:
        task_store[task_id] = {
            'status': 'failed',
            'error': str(e)
        }

四、Docker 部署

4.1 Dockerfile

# Dockerfile
FROM nvidia/cuda:11.8-cudnn8-runtime-ubuntu22.04

# 系统依赖
RUN apt-get update && apt-get install -y \
    python3.10 python3-pip \
    libgl1-mesa-glx libglib2.0-0 \
    && rm -rf /var/lib/apt/lists/*

WORKDIR /app

# Python 依赖
COPY requirements.txt .
RUN pip install --no-cache-dir -r requirements.txt

# 复制代码
COPY app/ ./app/
COPY pretrained_models/ ./pretrained_models/

# 暴露端口
EXPOSE 8000

# 启动服务
CMD ["uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000", "--workers", "1"]

4.2 requirements.txt

fastapi==0.104.1
uvicorn[standard]==0.24.0
python-multipart==0.0.6
opencv-python-headless==4.8.1.78
torch==2.0.1
torchvision==0.15.2
numpy==1.24.3
Pillow==10.1.0
pydantic==2.5.0
basicsr==1.4.2
realesrgan==0.3.0

4.3 docker-compose.yml

version: '3.8'

services:
  api:
    build: .
    ports:
      - "8000:8000"
    environment:
      - NVIDIA_VISIBLE_DEVICES=all
      - CUDA_VISIBLE_DEVICES=0
    volumes:
      - ./pretrained_models:/app/pretrained_models
      - ./output:/app/output
    deploy:
      resources:
        reservations:
          devices:
            - driver: nvidia
              count: 1
              capabilities: [gpu]
    restart: unless-stopped
    healthcheck:
      test: ["CMD", "curl", "-f", "http://localhost:8000/health"]
      interval: 30s
      timeout: 10s
      retries: 3

  nginx:
    image: nginx:alpine
    ports:
      - "443:443"
      - "80:80"
    volumes:
      - ./nginx.conf:/etc/nginx/nginx.conf
      - ./ssl:/etc/nginx/ssl
    depends_on:
      - api
    restart: unless-stopped

4.4 启动

# 构建并启动
docker-compose build
docker-compose up -d

# 查看日志
docker-compose logs -f api

# 测试
curl -X POST http://localhost:8000/restore \
  -F "file=@test_image.jpg" \
  -F "tasks=auto"

# 查看健康状态
curl http://localhost:8000/health

五、客户端调用示例

5.1 Python 客户端

import requests
import time


class RestorationClient:
    """图像复原 API 客户端"""
    def __init__(self, base_url="http://localhost:8000"):
        self.base_url = base_url

    def restore(self, image_path, tasks="auto", scale=4, noise_level=25):
        # 提交任务
        with open(image_path, 'rb') as f:
            response = requests.post(
                f"{self.base_url}/restore",
                files={"file": f},
                params={
                    "tasks": tasks,
                    "scale": scale,
                    "noise_level": noise_level
                }
            )
        data = response.json()
        task_id = data['task_id']
        print(f"任务提交: {task_id}")

        # 轮询结果
        while True:
            result = requests.get(
                f"{self.base_url}/task/{task_id}"
            ).json()

            if result['status'] == 'completed':
                # 下载结果
                img_response = requests.get(
                    f"{self.base_url}/download/{task_id}"
                )
                output_path = f"restored_{task_id}.png"
                with open(output_path, 'wb') as f:
                    f.write(img_response.content)
                print(f"完成! 处理时间: {result['processing_time']}ms")
                return output_path

            elif result['status'] == 'failed':
                raise Exception(f"处理失败: {result['error']}")

            time.sleep(1)


# 使用
client = RestorationClient("http://localhost:8000")
result_path = client.restore("my_image.jpg", tasks="denoise,super_resolution")

5.2 curl 调用

# 单任务
curl -X POST http://localhost:8000/restore \
  -F "file=@noisy_image.jpg" \
  -F 'tasks=denoise'

# 多任务串联
curl -X POST http://localhost:8000/restore \
  -F "file=@hazy_image.jpg" \
  -F 'tasks=dehaze,super_resolution' \
  -F 'scale=4'

# 自动模式
curl -X POST http://localhost:8000/restore \
  -F "file=@unknown_image.jpg" \
  -F 'tasks=auto'

六、生产环境 Checklist

项目 说明 状态
GPU 推理 使用 TensorRT/ONNX 加速 推荐
限流 Nginx rate limiting 必须
认证 API Key / JWT 推荐
日志 结构化日志 + ELK 推荐
监控 Prometheus + Grafana 推荐
队列 Redis + Celery(高并发) 推荐
缓存 相同图像重复请求缓存 可选
超时控制 大图像超时处理 必须

七、总结

本文从工程角度将 Real-ESRGAN、DehazeFormer 和 Restormer 整合为可部署的微服务。核心设计要点:

  1. 模型管理器:单例模式 + 懒加载,节省显存
  2. 管线编排:固定任务顺序(降噪→去模糊→去雾→超分)
  3. 异步处理:大图像不阻塞 API 响应
  4. Docker 部署:GPU 支持 + 健康检查 + 反向代理

这套方案可直接用于生产环境,支撑日均万级图像处理请求。

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐