Files
scan/api.py

371 lines
13 KiB
Python
Raw Normal View History

2026-09-14 06:09:13 +03:00
# /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:
2026-09-14 07:26:47 +03:00
DOCRES_MODEL_PATH = DOCREST_DIR / "data/weights" / "docres.pkl"
2026-09-14 06:09:13 +03:00
_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])
2026-09-14 06:09:13 +03:00
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])
2026-09-14 06:09:13 +03:00
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()
2026-09-14 06:09:13 +03:00
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
)
2026-09-14 06:09:13 +03:00
return Response(
content=result_jpeg,
media_type="image/jpeg",
headers={"Content-Disposition": 'inline; filename="unwrap.jpg"', "Cache-Control": "no-store"},
2026-09-14 06:09:13 +03:00
)
async def _docres_endpoint_async(request: Request, file: UploadFile, task: str):
payload = await file.read()
2026-09-14 06:09:13 +03:00
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
)
2026-09-14 06:09:13 +03:00
return Response(
content=result_jpeg,
media_type="image/jpeg",
headers={"Content-Disposition": f'inline; filename="{task}.jpg"', "Cache-Control": "no-store"},
2026-09-14 06:09:13 +03:00
)
@app.post("/deblur")
async def deblur(request: Request, file: UploadFile = File(...)):
return await _docres_endpoint_async(request, file, "deblur")
2026-09-14 06:09:13 +03:00
@app.post("/illumination_correct")
async def illumination_correct(request: Request, file: UploadFile = File(...)):
return await _docres_endpoint_async(request, file, "illumination_correct")
2026-09-14 06:09:13 +03:00
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host=API_HOST, port=API_PORT, workers=1)