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])
|
|
|
|
|
|
|
|
|
|
|
|
ok, buf = cv2.imencode(".png", result_bgr)
|
|
|
|
|
|
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
|
|
|
|
|
|
|
|
|
|
|
|
ok, buf = cv2.imencode(".png", np.ascontiguousarray(result))
|
|
|
|
|
|
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")
|
|
|
|
|
|
def unwrap(request: Request, file: UploadFile = File(...)):
|
|
|
|
|
|
try:
|
|
|
|
|
|
file.file.seek(0)
|
|
|
|
|
|
payload = file.file.read()
|
|
|
|
|
|
except Exception as exc:
|
|
|
|
|
|
raise HTTPException(status_code=400, detail="Cannot read uploaded file") from exc
|
|
|
|
|
|
if not payload:
|
|
|
|
|
|
raise HTTPException(status_code=400, detail="Uploaded file is empty")
|
|
|
|
|
|
|
|
|
|
|
|
result_png = process_unwrap(request.app.state.docscanner_net, payload)
|
|
|
|
|
|
return Response(
|
|
|
|
|
|
content=result_png,
|
|
|
|
|
|
media_type="image/png",
|
|
|
|
|
|
headers={"Content-Disposition": 'inline; filename="unwrap.png"', "Cache-Control": "no-store"},
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _docres_endpoint(request: Request, file: UploadFile, task: str):
|
|
|
|
|
|
try:
|
|
|
|
|
|
file.file.seek(0)
|
|
|
|
|
|
payload = file.file.read()
|
|
|
|
|
|
except Exception as exc:
|
|
|
|
|
|
raise HTTPException(status_code=400, detail="Cannot read uploaded file") from exc
|
|
|
|
|
|
if not payload:
|
|
|
|
|
|
raise HTTPException(status_code=400, detail="Uploaded file is empty")
|
|
|
|
|
|
|
|
|
|
|
|
result_png = process_docres(request.app.state.docres_model, payload, task)
|
|
|
|
|
|
return Response(
|
|
|
|
|
|
content=result_png,
|
|
|
|
|
|
media_type="image/png",
|
|
|
|
|
|
headers={"Content-Disposition": f'inline; filename="{task}.png"', "Cache-Control": "no-store"},
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@app.post("/deblur")
|
|
|
|
|
|
def deblur(request: Request, file: UploadFile = File(...)):
|
|
|
|
|
|
return _docres_endpoint(request, file, "deblur")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@app.post("/illumination_correct")
|
|
|
|
|
|
def illumination_correct(request: Request, file: UploadFile = File(...)):
|
|
|
|
|
|
return _docres_endpoint(request, file, "illumination_correct")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
if __name__ == "__main__":
|
|
|
|
|
|
import uvicorn
|
|
|
|
|
|
uvicorn.run(app, host=API_HOST, port=API_PORT, workers=1)
|