# /pkg/gnu/unitprint/backend/scan/api.py import os import sys import threading import io import warnings from contextlib import asynccontextmanager from pathlib import Path import cv2 import numpy as np import torch import torch.nn as nn import torch.nn.functional as F from PIL import Image, ImageOps from fastapi import FastAPI, File, HTTPException, Request, UploadFile from fastapi.responses import Response warnings.filterwarnings('ignore') # --- PATHS --- SCAN_DIR = Path(__file__).resolve().parent DOCSCANNER_DIR = SCAN_DIR / "DocScanner" DOCREST_DIR = SCAN_DIR / "DocRes" # Добавляем папки проектов в sys.path, чтобы Python мог найти их внутренние модули. # Импорты ниже сработают корректно, так как у проектов нет пересекающихся имен файлов # (DocScanner использует `model.py` и `seg.py`, DocRes использует `models/` и `utils.py`). for p in (DOCREST_DIR, DOCSCANNER_DIR): p_str = str(p) if p_str not in sys.path: sys.path.insert(0, p_str) # --- DOCSCANNER IMPORTS --- from model import DocScanner as DocScannerNet from seg import U2NETP # --- DOCRES IMPORTS --- from utils import convert_state_dict from models import restormer_arch from data.preprocess.crop_merge_image import stride_integral DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") USE_HALF = ( DEVICE.type == "cuda" and os.getenv("DOCRES_USE_HALF", "1").strip().lower() not in ("0", "false", "no") ) API_HOST = os.getenv("API_HOST", "127.0.0.1") API_PORT = int(os.getenv("API_PORT", "15200")) MAX_SIZE = max(1, int(os.getenv("DOCRES_MAX_SIZE", "1800"))) DOCSCANNER_SEG_PATH = DOCSCANNER_DIR / "model_pretrained" / "seg.pth" DOCSCANNER_BM_PATH = DOCSCANNER_DIR / "model_pretrained" / "DocScanner-L.pth" _docres_model_env = os.getenv("DOCRES_MODEL_PATH") if _docres_model_env: DOCRES_MODEL_PATH = Path(_docres_model_env.strip()).expanduser() if not DOCRES_MODEL_PATH.is_absolute(): DOCRES_MODEL_PATH = (SCAN_DIR / DOCRES_MODEL_PATH).resolve() else: DOCRES_MODEL_PATH = DOCREST_DIR / "data/weights" / "docres.pkl" _lock = threading.Lock() # ===================================================================== # --- DOCSCANNER LOGIC --- # ===================================================================== class DocScannerWrapper(nn.Module): def __init__(self): super().__init__() self.msk = U2NETP(3, 1) self.bm = DocScannerNet() def forward(self, x): msk, *_ = self.msk(x) bm = self.bm((msk > 0.5).float() * x, iters=12, test_mode=True) return (2 * (bm / 286.8) - 1) * 0.99 def load_docscanner_model(model, path, strip_prefix=False): state_dict = model.state_dict() try: pretrained = torch.load(path, map_location=DEVICE, weights_only=False) except TypeError: pretrained = torch.load(path, map_location=DEVICE) if strip_prefix: pretrained = {k[6:]: v for k, v in pretrained.items() if k[6:] in state_dict} else: pretrained = {k: v for k, v in pretrained.items() if k in state_dict} state_dict.update(pretrained) model.load_state_dict(state_dict) return model def process_unwrap(net, payload: bytes) -> bytes: try: img = Image.open(io.BytesIO(payload)) try: img = ImageOps.exif_transpose(img) except Exception: pass img = img.convert("RGB") except Exception as exc: raise HTTPException(status_code=400, detail="Uploaded file is not a valid image") from exc image = np.asarray(img, dtype=np.float32) / 255.0 height, width = image.shape[:2] tensor = ( torch.from_numpy(cv2.resize(image, (288, 288)).transpose(2, 0, 1)) .float() .unsqueeze(0) .to(DEVICE) ) with _lock, torch.no_grad(): bm = net(tensor).cpu() flow_x = cv2.blur(cv2.resize(bm[0, 0].numpy(), (width, height)), (3, 3)) flow_y = cv2.blur(cv2.resize(bm[0, 1].numpy(), (width, height)), (3, 3)) flow = torch.from_numpy(np.stack([flow_x, flow_y], axis=2)).unsqueeze(0).float() out = F.grid_sample( torch.from_numpy(image).permute(2, 0, 1).unsqueeze(0).float(), flow, align_corners=True, ) result = (out[0] * 255).permute(1, 2, 0).numpy() result = np.clip(result, 0, 255).astype(np.uint8) result_bgr = np.ascontiguousarray(result[:, :, ::-1]) # ИЗМЕНЕНО: кодирование в JPEG с качеством 100% ok, buf = cv2.imencode(".jpg", result_bgr, [int(cv2.IMWRITE_JPEG_QUALITY), 100]) if not ok: raise HTTPException(status_code=500, detail="Failed to encode result image") return buf.tobytes() # ===================================================================== # --- DOCRES LOGIC --- # ===================================================================== def load_docres_model(model, path): try: ckpt = torch.load(path, map_location=DEVICE, weights_only=False) except TypeError: ckpt = torch.load(path, map_location=DEVICE) if isinstance(ckpt, dict): state = None for key in ("model_state", "state_dict", "model", "params", "params_ema"): if key in ckpt and isinstance(ckpt[key], dict): state = ckpt[key] break if state is None: state = ckpt else: raise RuntimeError(f"Unsupported checkpoint format: {path}") state = convert_state_dict(state) model.load_state_dict(state) model.eval() model.to(DEVICE) if USE_HALF: model.half() else: model.float() return model def resize_if_needed(img: np.ndarray, max_size: int = MAX_SIZE) -> np.ndarray: h, w = img.shape[:2] longest = max(h, w) if longest <= 0: raise HTTPException(status_code=400, detail="Image has zero size") scale = min(max_size / longest, 1.0) if scale < 1.0: new_w = max(1, int(w * scale)) new_h = max(1, int(h * scale)) return cv2.resize(img, (new_w, new_h), interpolation=cv2.INTER_AREA) return img def deblur_prompt(img: np.ndarray) -> np.ndarray: x = cv2.Sobel(img, cv2.CV_16S, 1, 0) y = cv2.Sobel(img, cv2.CV_16S, 0, 1) absX = cv2.convertScaleAbs(x) absY = cv2.convertScaleAbs(y) high_frequency = cv2.addWeighted(absX, 0.5, absY, 0.5, 0) high_frequency = cv2.cvtColor(high_frequency, cv2.COLOR_BGR2GRAY) high_frequency = cv2.cvtColor(high_frequency, cv2.COLOR_GRAY2BGR) return high_frequency def appearance_prompt(img: np.ndarray) -> np.ndarray: h, w = img.shape[:2] img_1024 = cv2.resize(img, (1024, 1024)) planes = cv2.split(img_1024) norm_planes = [] for plane in planes: dilated_img = cv2.dilate(plane, np.ones((7, 7), np.uint8)) bg_img = cv2.medianBlur(dilated_img, 21) diff_img = 255 - cv2.absdiff(plane, bg_img) norm_img = cv2.normalize(diff_img, None, alpha=0, beta=255, norm_type=cv2.NORM_MINMAX, dtype=cv2.CV_8UC1) norm_planes.append(norm_img) result_norm = cv2.merge(norm_planes) result_norm = cv2.resize(result_norm, (w, h)) return result_norm def _run_restormer(model, in_im: np.ndarray, padding_h: int, padding_w: int) -> np.ndarray: arr = in_im.astype(np.float32) / 255.0 with _lock, torch.no_grad(): tensor = torch.from_numpy(arr.transpose(2, 0, 1)).unsqueeze(0).to(DEVICE) tensor = tensor.half() if USE_HALF else tensor.float() pred = model(tensor) pred = torch.clamp(pred, 0, 1) pred = pred[0].permute(1, 2, 0).cpu().numpy() pred = (pred * 255).astype(np.uint8) out = pred[padding_h:, padding_w:] if out.size == 0: raise HTTPException(status_code=500, detail="Empty result after padding crop") return out def deblur_image(model, img: np.ndarray) -> np.ndarray: img = resize_if_needed(img) in_im, padding_h, padding_w = stride_integral(img, 8) prompt = deblur_prompt(in_im) in_im = np.concatenate((in_im, prompt), -1) return _run_restormer(model, in_im, padding_h, padding_w) def appearance_image(model, img: np.ndarray) -> np.ndarray: img = resize_if_needed(img) prompt = appearance_prompt(img) in_im = np.concatenate((img, prompt), -1) in_im, padding_h, padding_w = stride_integral(in_im, 8) return _run_restormer(model, in_im, padding_h, padding_w) def process_docres(model, payload: bytes, task: str) -> bytes: img = cv2.imdecode(np.frombuffer(payload, np.uint8), cv2.IMREAD_COLOR) if img is None or img.size == 0: raise HTTPException(status_code=400, detail="Uploaded file is not a valid image") try: if task == "deblur": result = deblur_image(model, img) elif task == "illumination_correct": result = appearance_image(model, img) else: raise HTTPException(status_code=500, detail=f"Unknown task: {task}") except HTTPException: raise except Exception as exc: raise HTTPException(status_code=500, detail=f"{task} failed: {exc}") from exc # ИЗМЕНЕНО: кодирование в JPEG с качеством 100% ok, buf = cv2.imencode(".jpg", np.ascontiguousarray(result), [int(cv2.IMWRITE_JPEG_QUALITY), 100]) if not ok: raise HTTPException(status_code=500, detail="Failed to encode result image") return buf.tobytes() # ===================================================================== # --- FASTAPI APP --- # ===================================================================== @asynccontextmanager async def lifespan(app: FastAPI): if not DOCSCANNER_SEG_PATH.exists(): raise RuntimeError(f"DocScanner seg.pth not found: {DOCSCANNER_SEG_PATH}") if not DOCSCANNER_BM_PATH.exists(): raise RuntimeError(f"DocScanner DocScanner-L.pth not found: {DOCSCANNER_BM_PATH}") if not DOCRES_MODEL_PATH.exists(): raise RuntimeError(f"DocRes checkpoint not found: {DOCRES_MODEL_PATH}") # Init DocScanner docscanner_net = DocScannerWrapper().to(DEVICE).eval() load_docscanner_model(docscanner_net.msk, DOCSCANNER_SEG_PATH, strip_prefix=True) load_docscanner_model(docscanner_net.bm, DOCSCANNER_BM_PATH, strip_prefix=False) # Init DocRes docres_model = restormer_arch.Restormer( inp_channels=6, out_channels=3, dim=48, num_blocks=[2, 3, 3, 4], num_refinement_blocks=4, heads=[1, 2, 4, 8], ffn_expansion_factor=2.66, bias=False, LayerNorm_type="WithBias", dual_pixel_task=True, ) load_docres_model(docres_model, DOCRES_MODEL_PATH) app.state.docscanner_net = docscanner_net app.state.docres_model = docres_model yield app = FastAPI(title="Scan Backend API", lifespan=lifespan) @app.get("/health") def health(): return { "status": "ok", "device": str(DEVICE), "docres_half_precision": USE_HALF, "docres_model_path": str(DOCRES_MODEL_PATH), "docscanner_paths": { "seg": str(DOCSCANNER_SEG_PATH), "bm": str(DOCSCANNER_BM_PATH) } } @app.post("/unwrap") async def unwrap(request: Request, file: UploadFile = File(...)): # Читаем асинхронно — Starlette сам корректно обработает SpooledTemporaryFile payload = await file.read() if not payload: raise HTTPException(status_code=400, detail="Uploaded file is empty") # Тяжёлую работу отправляем в threadpool, чтобы не блокировать event loop import asyncio loop = asyncio.get_running_loop() result_jpeg = await loop.run_in_executor( None, process_unwrap, request.app.state.docscanner_net, payload ) return Response( content=result_jpeg, media_type="image/jpeg", headers={"Content-Disposition": 'inline; filename="unwrap.jpg"', "Cache-Control": "no-store"}, ) async def _docres_endpoint_async(request: Request, file: UploadFile, task: str): payload = await file.read() if not payload: raise HTTPException(status_code=400, detail="Uploaded file is empty") import asyncio loop = asyncio.get_running_loop() result_jpeg = await loop.run_in_executor( None, process_docres, request.app.state.docres_model, payload, task ) return Response( content=result_jpeg, media_type="image/jpeg", headers={"Content-Disposition": f'inline; filename="{task}.jpg"', "Cache-Control": "no-store"}, ) @app.post("/deblur") async def deblur(request: Request, file: UploadFile = File(...)): return await _docres_endpoint_async(request, file, "deblur") @app.post("/illumination_correct") async def illumination_correct(request: Request, file: UploadFile = File(...)): return await _docres_endpoint_async(request, file, "illumination_correct") if __name__ == "__main__": import uvicorn uvicorn.run(app, host=API_HOST, port=API_PORT, workers=1)