File size: 5,177 Bytes
fbcc547
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
#!/usr/bin/env python3
"""Oracle onnxruntime pour la parité PP-OCRv6 tiny (burn vs ORT).

Deux modes par réseau :
  * « exact » : les tenseurs d'entrée pré-traités par burn_ppocr (fichiers .f32
    little-endian + formes dans le manifeste) sont donnés à ORT tels quels →
    ne compare que les réseaux ;
  * « own » : ce script refait le pré-traitement de PaddleOCR lui-même (PIL
    BILINEAR, BGR, mean/std, crops à partir des boîtes JSON de burn_ppocr) →
    compare réseaux + pré-traitement.

Usage :
  ppocr_ref.py --image hello.png --det det_pads.onnx --rec rec_pads.onnx \
      --work DIR   # DIR contient boxes.json (+ det_input.f32, rec_input.f32, manifest.json)
Écrit dans DIR : det_out_exact.f32, det_out_own.f32, rec_out_exact.f32,
rec_out_own.f32 et ref.json (formes).
"""
import argparse
import json
import math
import os

import numpy as np
import onnxruntime as ort
from PIL import Image

MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32)
STD = np.array([0.229, 0.224, 0.225], dtype=np.float32)


def det_resize_dims(w, h, limit, max_side):
    ratio = limit / min(w, h) if min(w, h) < limit else 1.0
    if max(w, h) * ratio > max_side:
        ratio = max_side / max(w, h)
    rw, rh = int(w * ratio), int(h * ratio)
    rw = max(int(round(rw / 32) * 32), 32)
    rh = max(int(round(rh / 32) * 32), 32)
    return rw, rh


def det_own_input(img, limit, max_side):
    rw, rh = det_resize_dims(img.width, img.height, limit, max_side)
    resized = np.asarray(img.resize((rw, rh), Image.BILINEAR), dtype=np.float32)  # HWC RGB
    bgr = resized[:, :, ::-1] / 255.0
    norm = (bgr - MEAN) / STD
    return norm.transpose(2, 0, 1)[None].astype(np.float32)


def crop_rotate(arr, x0, y0, x1, y1):
    c = arr[y0:y1, x0:x1]
    h, w = c.shape[:2]
    if w > 0 and h / w >= 1.5:
        c = np.rot90(c)
    return c


def rec_own_input(img, boxes, rec_h, img_w):
    arr = np.asarray(img)
    crops = [crop_rotate(arr, *b) for b in boxes]
    out = np.zeros((len(crops), 3, rec_h, img_w), dtype=np.float32)
    for i, c in enumerate(crops):
        h, w = c.shape[:2]
        if h == 0 or w == 0:
            continue
        rw = min(max(int(math.ceil(rec_h * w / h)), 1), img_w)
        r = np.asarray(Image.fromarray(c).resize((rw, rec_h), Image.BILINEAR), dtype=np.float32)
        bgr = r[:, :, ::-1] / 255.0
        norm = (bgr - 0.5) / 0.5
        out[i, :, :, :rw] = norm.transpose(2, 0, 1)
    return out


def read_f32(path, shape):
    return np.fromfile(path, dtype="<f4").reshape(shape)


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--image", required=True)
    ap.add_argument("--det", required=True)
    ap.add_argument("--rec", required=True)
    ap.add_argument("--work", required=True)
    ap.add_argument("--limit", type=int, default=736)
    ap.add_argument("--max-side", type=int, default=4000)
    ap.add_argument("--rec-height", type=int, default=48)
    ap.add_argument("--rec-width", type=int, default=320, help="largeur du lot calculée par burn_ppocr")
    a = ap.parse_args()

    so = ort.SessionOptions()
    so.intra_op_num_threads = 1
    det = ort.InferenceSession(a.det, so, providers=["CPUExecutionProvider"])
    rec = ort.InferenceSession(a.rec, so, providers=["CPUExecutionProvider"])
    din, rin = det.get_inputs()[0].name, rec.get_inputs()[0].name

    img = Image.open(a.image).convert("RGB")
    work = a.work
    manifest = json.load(open(os.path.join(work, "manifest.json")))
    boxes = manifest["boxes"]
    out = {}

    # det, exact
    shape = manifest["det_input_shape"]
    x = read_f32(os.path.join(work, "det_input.f32"), shape)
    y = det.run(None, {din: x})[0]
    y.astype("<f4").tofile(os.path.join(work, "det_out_exact.f32"))
    out["det_exact_shape"] = list(y.shape)
    # det, own preprocessing
    x2 = det_own_input(img, a.limit, a.max_side)
    out["det_own_input_max_abs_diff"] = float(np.abs(x2 - x).max()) if x2.shape == x.shape else None
    y2 = det.run(None, {din: x2})[0]
    y2.astype("<f4").tofile(os.path.join(work, "det_out_own.f32"))
    out["det_own_shape"] = list(y2.shape)

    # rec, exact
    shape = manifest["rec_input_shape"]
    if shape[0] > 0:
        xr = read_f32(os.path.join(work, "rec_input.f32"), shape)
        yr = rec.run(None, {rin: xr})[0]
        yr.astype("<f4").tofile(os.path.join(work, "rec_out_exact.f32"))
        out["rec_exact_shape"] = list(yr.shape)
        xr2 = rec_own_input(img, boxes, a.rec_height, a.rec_width)
        out["rec_own_input_max_abs_diff"] = float(np.abs(xr2 - xr).max()) if xr2.shape == xr.shape else None
        yr2 = rec.run(None, {rin: xr2})[0]
        yr2.astype("<f4").tofile(os.path.join(work, "rec_out_own.f32"))
        out["rec_own_shape"] = list(yr2.shape)
        # texte décodé côté ORT (glouton, blank 0, espace dernier) pour lecture humaine
        idx = yr.argmax(axis=2)
        out["rec_exact_argmax"] = idx.tolist()
    else:
        out["rec_exact_shape"] = [0]

    json.dump(out, open(os.path.join(work, "ref.json"), "w"))
    print(json.dumps({k: v for k, v in out.items() if k != "rec_exact_argmax"}, indent=1))


if __name__ == "__main__":
    main()