inference переписан под CLI, удаление не рантайм файлов

This commit is contained in:
2026-09-12 03:25:38 +03:00
parent 7c854e418d
commit d824dd25e1
8 changed files with 61 additions and 361 deletions

1
.gitignore vendored Normal file
View File

@ -0,0 +1 @@
__pycache__

View File

@ -1,43 +0,0 @@
def Levenshtein_Distance(str1, str2):
matrix = [[ i + j for j in range(len(str2) + 1)] for i in range(len(str1) + 1)]
for i in range(1, len(str1)+1):
for j in range(1, len(str2)+1):
if(str1[i-1] == str2[j-1]):
d = 0
else:
d = 1
matrix[i][j] = min(matrix[i-1][j]+1, matrix[i][j-1]+1, matrix[i-1][j-1]+d)
return matrix[len(str1)][len(str2)]
def cal_cer_ed(path_ours, tail='_rec'):
path_gt='./GT/'
N=66
cer1=[]
cer2=[]
ed1=[]
ed2=[]
check=[0 for _ in range(N+1)]
lis=[1,2,3,4,5,6,7,9,10,21,22,23,24,27,30,31,32,36,38,40,41,44,45,46,47,48,50,51,52,53] # DocTr (Setting 1)
# lis=[1,9,10,12,19,20,21,22,23,24,30,31,32,34,35,36,37,38,39,40,44,45,46,47,49] # DewarpNet (Setting 2)
for i in range(1,N):
if i not in lis:
continue
gt=Image.open(path_gt+str(i)+'.png')
img1=Image.open(path_ours+str(i)+'_1' + tail)
img2=Image.open(path_ours+str(i)+'_2' + tail)
content_gt=pytesseract.image_to_string(gt)
content1=pytesseract.image_to_string(img1)
content2=pytesseract.image_to_string(img2)
l1=Levenshtein_Distance(content_gt,content1)
l2=Levenshtein_Distance(content_gt,content2)
ed1.append(l1)
ed2.append(l2)
cer1.append(l1/len(content_gt))
cer2.append(l2/len(content_gt))
check[i]=cer1[-1]
print('CER: ', (np.mean(cer1)+np.mean(cer2)) / 2.)
print('ED: ', (np.mean(ed1)+np.mean(ed2)) / 2.)
def evalu(path_ours, tail):
cal_cer_ed(path_ours, tail)

64
eval.m
View File

@ -1,64 +0,0 @@
path_rec = "xxx"; % rectified image path
path_scan = './scan/'; % scan image path
label_path = './layout/'; % layout result path
tarea = 598400;
ms1 = 0;
ld1 = 0;
lid1 = 0;
ms2 = 0;
ld2 = 0;
lid2 = 0;
wv = 0;
wh = 0;
sprintf(path_rec)
for i=1:65
path_rec_1 = sprintf("%s%d%s", path_rec, i, '_1 copy_rec.png'); % rectified image path
path_rec_2 = sprintf("%s%d%s", path_rec, i, '_2 copy_rec.png'); % rectified image path
path_scan_new = sprintf("%s%d%s", path_scan, i, '.png'); % corresponding scan image path
bbox_i_path = sprintf("%s%d%s", label_path, i, '.txt'); % corresponding layout txt path
% imread and rgb2gray
A1 = imread(path_rec_1);
A2 = imread(path_rec_2);
% if i == 64
% A1 = rot90(A1,-2);
% A2 = rot90(A2,-2);
% end
ref = imread(path_scan_new);
A1 = rgb2gray(A1);
A2 = rgb2gray(A2);
ref = rgb2gray(ref);
bbox_i = read_txt(bbox_i_path);
bbox_i = bbox_i + 1; % python index starts from 0
% resize
b = sqrt(tarea/size(ref,1)/size(ref,2));
ref = imresize(ref,b);
A1 = imresize(A1,[size(ref,1),size(ref,2)]);
A2 = imresize(A2,[size(ref,1),size(ref,2)]);
scaled_bbox_i = bbox_i * b * 0.5;
scaled_bbox_i = round(scaled_bbox_i);
scaled_bbox_i = max(scaled_bbox_i, 1);
% calculate
[ms_1, ld_1, lid_1, W_v_1, W_h_1] = evalUnwarp(A1, ref, scaled_bbox_i);
[ms_2, ld_2, lid_2, W_v_2, W_h_2] = evalUnwarp(A2, ref, scaled_bbox_i);
ms1 = ms1 + ms_1;
ms2 = ms2 + ms_2;
ld1 = ld1 + ld_1;
ld2 = ld2 + ld_2;
lid1 = lid1 + lid_1;
lid2 = lid2 + lid_2;
wv = wv + W_v_1 + W_v_2;
wh = wh + W_h_1 + W_h_2;
end
ms = (ms1 + ms2) / 130 % MS-SSIM
ld = (ld1 + ld2) / 130 % local distortion
li_d = (lid1 + lid2) / 130 % line distortion
wv = wv / 130 % wv index
wh = wh / 130 % wh index

View File

@ -1,102 +0,0 @@
function [ms, ld, li_d, wv, wh] = evalUnwarp(A, ref, data)
%EVALUNWARP compute MSSSIM and LD between the unwarped image and the scan
% A: unwarped image
% ref: reference image, the scan image
% ms: returned MS-SSIM value
% ld: returned local distortion value
% Matlab image processing toolbox is necessary to compute ssim. The weights
% for multi-scale ssim is directly adopted from:
%
% Wang, Zhou, Eero P. Simoncelli, and Alan C. Bovik. "Multiscale structural
% similarity for image quality assessment." In Signals, Systems and Computers,
% 2004. Conference Record of the Thirty-Seventh Asilomar Conference on, 2003.
%
% Local distortion relies on the paper:
% Liu, Ce, Jenny Yuen, and Antonio Torralba. "Sift flow: Dense correspondence
% across scenes and its applications." In PAMI, 2010.
%
% and its implementation:
% https://people.csail.mit.edu/celiu/SIFTflow/
x = A;
y = ref;
im1=imresize(imfilter(y,fspecial('gaussian',7,1.),'same','replicate'),0.5,'bicubic');
im2=imresize(imfilter(x,fspecial('gaussian',7,1.),'same','replicate'),0.5,'bicubic');
im1=im2double(im1);
im2=im2double(im2);
cellsize=3;
gridspacing=1;
sift1 = mexDenseSIFT(im1,cellsize,gridspacing);
sift2 = mexDenseSIFT(im2,cellsize,gridspacing);
SIFTflowpara.alpha=2*255;
SIFTflowpara.d=40*255;
SIFTflowpara.gamma=0.005*255;
SIFTflowpara.nlevels=4;
SIFTflowpara.wsize=2;
SIFTflowpara.topwsize=10;
SIFTflowpara.nTopIterations = 60;
SIFTflowpara.nIterations= 30;
[vx,vy,~]=SIFTflowc2f(sift1,sift2,SIFTflowpara);
rows1p = size(im1,1);
cols1p = size(im1,2);
% Li-D
rowstd_sum = 0;
for i = 1:rows1p
rowstd = std(vy(i, :),1);
rowstd_sum = rowstd_sum + rowstd;
end
rowstd_mean = rowstd_sum / rows1p;
colstd_sum = 0;
for i = 1:cols1p
colstd = std(vx(:, i),1);
colstd_sum = colstd_sum + colstd;
end
colstd_mean = colstd_sum / cols1p;
li_d = (rowstd_mean + colstd_mean) / 2;
% LD
d = sqrt(vx.^2 + vy.^2);
ld = mean(d(:));
% MS-SSIM
wt = [0.0448 0.2856 0.3001 0.2363 0.1333];
ss = zeros(5, 1);
for s = 1 : 5
ss(s) = ssim(x, y);
x = impyramid(x, 'reduce');
y = impyramid(y, 'reduce');
end
ms = wt * ss;
% wv and wh
rowstd_sum = 0;
for i = 1:size(data, 1)
rowstd_top = std(vy(data(i,2), data(i,1):data(i,3)),1) / (data(i,3)-data(i,1));
rowstd_bot = std(vy(data(i,4), data(i,1):data(i,3)),1) / (data(i,3)-data(i,1));
rowstd_sum = rowstd_sum + rowstd_top + rowstd_bot;
end
wv = rowstd_sum / (2 * size(data, 1));
colstd_sum = 0;
for i = 1:size(data, 1)
colstd_left = std(vx(data(i,2):data(i,4), data(i,1)),1) / (data(i,4)- data(i,2));
colstd_right = std(vx(data(i,2):data(i,4), data(i,3)),1) / (data(i,4)- data(i,2));
colstd_sum = colstd_sum + colstd_left + colstd_right;
end
wh = colstd_sum / (2 * size(data, 1));
end

View File

@ -1,115 +1,85 @@
import os, sys, argparse, warnings
import torch, torch.nn as nn, torch.nn.functional as F
import numpy as np, cv2
from PIL import Image
from model import DocScanner from model import DocScanner
from seg import U2NETP from seg import U2NETP
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
import cv2
import os
from PIL import Image
import argparse
import warnings
warnings.filterwarnings('ignore') warnings.filterwarnings('ignore')
# Запоминаем CWD пользователя ДО смены директории
USER_CWD = os.getcwd()
SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__))
os.chdir(SCRIPT_DIR)
class Net(nn.Module): class Net(nn.Module):
def __init__(self): def __init__(self):
super(Net, self).__init__() super().__init__()
self.msk = U2NETP(3, 1) self.msk = U2NETP(3, 1)
self.bm = DocScanner() # 矫正 self.bm = DocScanner()
def forward(self, x): def forward(self, x):
msk, _1,_2,_3,_4,_5,_6 = self.msk(x) msk, *_ = self.msk(x)
msk = (msk > 0.5).float() bm = self.bm((msk > 0.5).float() * x, iters=12, test_mode=True)
x = msk * x return (2 * (bm / 286.8) - 1) * 0.99
bm = self.bm(x, iters=12, test_mode=True)
bm = (2 * (bm / 286.8) - 1) * 0.99
return bm
def reload_seg_model(model, path=""): def load_model(model, path, strip_prefix=False):
if not bool(path): state_dict = model.state_dict()
return model pretrained = torch.load(path, map_location='cuda:0')
if strip_prefix:
pretrained = {k[6:]: v for k, v in pretrained.items() if k[6:] in state_dict}
else: else:
model_dict = model.state_dict() pretrained = {k: v for k, v in pretrained.items() if k in state_dict}
pretrained_dict = torch.load(path, map_location='cuda:0') state_dict.update(pretrained)
pretrained_dict = {k[6:]: v for k, v in pretrained_dict.items() if k[6:] in model_dict} model.load_state_dict(state_dict)
model_dict.update(pretrained_dict) return model
model.load_state_dict(model_dict)
return model
def reload_rec_model(model, path=""): def resolve_user_path(path):
if not bool(path): """Если путь относительный — считаем его относительно CWD пользователя, не скрипта"""
return model if os.path.isabs(path):
else: return path
model_dict = model.state_dict() return os.path.join(USER_CWD, path)
pretrained_dict = torch.load(path, map_location='cuda:0')
pretrained_dict = {k: v for k, v in pretrained_dict.items() if k in model_dict}
model_dict.update(pretrained_dict)
model.load_state_dict(model_dict)
return model
def rec(seg_model_path, rec_model_path, distorrted_path, save_path):
# distorted images list
img_list = os.listdir(distorrted_path)
# creat save path for rectified images
if not os.path.exists(save_path):
os.makedirs(save_path)
# net init
net = Net().cuda()
# reload seg model
reload_seg_model(net.msk, seg_model_path)
# reload rec model
reload_rec_model(net.bm, rec_model_path)
net.eval()
for img_path in img_list:
name = img_path.split('.')[-2] # image name
img_path = distorrted_path + img_path # image path
im_ori = np.array(Image.open(img_path))[:, :, :3] / 255.
h, w, _ = im_ori.shape
im = cv2.resize(im_ori, (288, 288))
im = im.transpose(2, 0, 1)
im = torch.from_numpy(im).float().unsqueeze(0)
with torch.no_grad():
bm = net(im.cuda())
bm = bm.cpu()
# save rectified image
bm0 = cv2.resize(bm[0, 0].numpy(), (w, h)) # x flow
bm1 = cv2.resize(bm[0, 1].numpy(), (w, h)) # y flow
bm0 = cv2.blur(bm0, (3, 3))
bm1 = cv2.blur(bm1, (3, 3))
lbl = torch.from_numpy(np.stack([bm0, bm1], axis=2)).unsqueeze(0) # h * w * 2
out = F.grid_sample(torch.from_numpy(im_ori).permute(2, 0, 1).unsqueeze(0).float(), lbl, align_corners=True)
cv2.imwrite(save_path + name + '_rec' + '.png', (((out[0]*255).permute(1, 2, 0).numpy())[:,:,::-1]).astype(np.uint8))
def main(): def main():
parser = argparse.ArgumentParser() parser = argparse.ArgumentParser()
parser.add_argument('--seg_model_path', default='./model_pretrained/seg.pth') parser.add_argument('-i', '--input', required=True)
parser.add_argument('--rec_model_path', default='./model_pretrained/DocScanner-L.pth') parser.add_argument('-o', '--output', required=True)
parser.add_argument('--distorrted_path', default='./distorted/')
parser.add_argument('--rectified_path', default='./rectified/')
opt = parser.parse_args() opt = parser.parse_args()
rec(seg_model_path=opt.seg_model_path, input_path = resolve_user_path(opt.input)
rec_model_path=opt.rec_model_path, output_path = resolve_user_path(opt.output)
distorrted_path=opt.distorrted_path,
save_path=opt.rectified_path) net = Net().cuda().eval()
load_model(net.msk, f'{SCRIPT_DIR}/model_pretrained/seg.pth', strip_prefix=True)
load_model(net.bm, f'{SCRIPT_DIR}/model_pretrained/DocScanner-L.pth')
image = np.array(Image.open(input_path))[:, :, :3] / 255.0
height, width = image.shape[:2]
tensor = torch.from_numpy(cv2.resize(image, (288, 288)).transpose(2, 0, 1)).float().unsqueeze(0)
with torch.no_grad():
bm = net(tensor.cuda()).cpu()
flow = torch.from_numpy(np.stack([
cv2.blur(cv2.resize(bm[0, 0].numpy(), (width, height)), (3, 3)),
cv2.blur(cv2.resize(bm[0, 1].numpy(), (width, height)), (3, 3))
], axis=2)).unsqueeze(0)
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()[:, :, ::-1]).astype(np.uint8)
out_dir = os.path.dirname(output_path)
if out_dir:
os.makedirs(out_dir, exist_ok=True)
cv2.imwrite(output_path, result)
print(f"[OK] {output_path}")
if __name__ == "__main__": if __name__ == "__main__":

Binary file not shown.

BIN
model_pretrained/seg.pth Normal file

Binary file not shown.

View File

@ -1,62 +0,0 @@
The images for OCR evaluation of DocUNet Benchmark.
# Setting 1 (Setting from DocTr)
# Total 30 * 2 = 60 images.
./scan/1.png
./scan/2.png
./scan/3.png
./scan/4.png
./scan/5.png
./scan/6.png
./scan/7.png
./scan/9.png
./scan/10.png
./scan/21.png
./scan/22.png
./scan/23.png
./scan/24.png
./scan/27.png
./scan/30.png
./scan/31.png
./scan/32.png
./scan/36.png
./scan/38.png
./scan/40.png
./scan/41.png
./scan/44.png
./scan/45.png
./scan/46.png
./scan/47.png
./scan/48.png
./scan/50.png
./scan/51.png
./scan/52.png
./scan/53.png
# Setting 2 (Setting from DewarpNet)
# Link: https://github.com/cvlab-stonybrook/DewarpNet/blob/master/eval/ocr_eval/ocr_files.txt
# Total 25 * 2 = 50 images.
./scan/1.png
./scan/9.png
./scan/10.png
./scan/12.png
./scan/19.png
./scan/20.png
./scan/21.png
./scan/22.png
./scan/23.png
./scan/24.png
./scan/30.png
./scan/31.png
./scan/32.png
./scan/34.png
./scan/35.png
./scan/36.png
./scan/37.png
./scan/38.png
./scan/39.png
./scan/40.png
./scan/44.png
./scan/45.png
./scan/46.png
./scan/47.png
./scan/49.png