PPOCR_v6 / run_rec_onnx.py
philcuriosity1024's picture
Upload folder using huggingface_hub
6d39ea5 verified
Raw
History Blame Contribute Delete
17 kB
#!/usr/bin/env python3
"""
PP-OCRv6 Recognition ONNX Inference & Evaluation (standalone, zero Paddle dependency)
Dependencies:
numpy, opencv-python, onnxruntime, pyyaml
Usage:
# Single image inference
python ppocrv6_rec_onnx.py --rec_onnx rec.onnx --char_dict inference.yml --image crop.png
# Batch evaluation
python ppocrv6_rec_onnx.py --rec_onnx rec.onnx --char_dict inference.yml \\
--label_file labels.txt --dataset_root ./crops/ \\
--batch_size 8 --output_json result.json --verbose
"""
import argparse
import math
import os
from typing import List, Optional, Tuple, Union
import cv2
import numpy as np
import onnxruntime as ort
import yaml
# ============================================================================
# 1. Utilities
# ============================================================================
def _edit_distance(pred: str, target: str) -> Tuple[int, float]:
"""Compute Levenshtein edit distance (pure Python, no extra deps).
Returns:
(distance, normalized_distance) where normalized ∈ [0, 1].
"""
m, n = len(pred), len(target)
if m == 0:
return n, 1.0
if n == 0:
return m, 1.0
dp = list(range(n + 1))
for i in range(1, m + 1):
prev = dp[0]
dp[0] = i
for j in range(1, n + 1):
temp = dp[j]
if pred[i - 1] == target[j - 1]:
dp[j] = prev
else:
dp[j] = 1 + min(prev, dp[j], dp[j - 1])
prev = temp
distance = dp[n]
normalized = distance / max(m, n)
return distance, normalized
def _load_char_dict(source: Union[str, List[str]]) -> List[str]:
"""Load character dictionary from .yml, .txt, or list."""
if isinstance(source, list):
return list(source)
ext = os.path.splitext(source)[1].lower()
if ext in (".yml", ".yaml"):
with open(source, "r", encoding="utf-8") as f:
cfg = yaml.safe_load(f)
dic = cfg.get("PostProcess", {}).get("character_dict", [])
if not dic:
raise ValueError(f"No PostProcess.character_dict found in {source}")
return dic
elif ext == ".txt":
with open(source, "r", encoding="utf-8") as f:
return [line.strip("\n\r") for line in f.readlines()]
else:
raise ValueError(
f"Unsupported char_dict source: {source}. Use .yml, .txt, or list."
)
def _resize_norm_img(
img: np.ndarray,
image_shape: Tuple[int, int, int] = (3, 48, 320),
max_wh_ratio: Optional[float] = None,
) -> np.ndarray:
"""Resize and normalize a cropped text image for recognition.
Args:
img: BGR crop image (H, W, 3).
image_shape: (C, H, W) target shape.
max_wh_ratio: precomputed max width/height ratio for batch. If None, derived from img.
"""
imgC, imgH, imgW = image_shape
if max_wh_ratio is None:
max_wh_ratio = imgW * 1.0 / imgH
h, w = img.shape[:2]
ratio = w * 1.0 / h
max_wh_ratio = max(max_wh_ratio, ratio)
max_wh_ratio = min(max_wh_ratio, imgW / imgH)
target_w = int(imgH * max_wh_ratio)
h, w = img.shape[:2]
ratio = w * 1.0 / h
if math.ceil(imgH * ratio) > target_w:
resized_w = target_w
else:
resized_w = int(math.ceil(imgH * ratio))
resized = cv2.resize(img, (resized_w, imgH))
resized = resized.astype("float32")
resized = resized.transpose((2, 0, 1)) / 255.0
resized -= 0.5
resized /= 0.5
padded = np.zeros((imgC, imgH, target_w), dtype=np.float32)
padded[:, :, 0:resized_w] = resized
return padded
# ============================================================================
# 2. CTC Decoder
# ============================================================================
class _CTCLabelDecode:
"""CTC greedy decoder for recognition output."""
def __init__(self, character_list: List[str], use_space_char: bool = True):
self.character_str = list(character_list)
if use_space_char:
self.character_str.append(" ")
dict_character = ["blank"] + self.character_str
self.character = dict_character
self.dict = {char: i for i, char in enumerate(dict_character)}
def decode(
self,
text_index: np.ndarray,
text_prob: Optional[np.ndarray] = None,
is_remove_duplicate: bool = True,
) -> List[Tuple[str, float]]:
result_list = []
batch_size = len(text_index)
for batch_idx in range(batch_size):
selection = np.ones(len(text_index[batch_idx]), dtype=bool)
if is_remove_duplicate:
selection[1:] = text_index[batch_idx][1:] != text_index[batch_idx][:-1]
selection &= text_index[batch_idx] != 0 # ignore blank
char_list = [
self.character[int(tid)]
for tid in text_index[batch_idx][selection]
]
if text_prob is not None:
conf_list = text_prob[batch_idx][selection]
else:
conf_list = np.ones(len(selection), dtype=np.float32)
if len(conf_list) == 0:
conf_list = np.array([0.0], dtype=np.float32)
text = "".join(char_list)
result_list.append((text, float(np.mean(conf_list))))
return result_list
def __call__(self, preds: np.ndarray) -> List[Tuple[str, float]]:
preds_idx = preds.argmax(axis=2)
preds_prob = preds.max(axis=2)
return self.decode(preds_idx, preds_prob, is_remove_duplicate=True)
# ============================================================================
# 3. Recognition Engine
# ============================================================================
class PPOCRv6RecOnnx:
def __init__(
self,
rec_onnx: str,
char_dict: Union[str, List[str]],
rec_image_shape: Tuple[int, int, int] = (3, 48, 320),
rec_batch_num: int = 6,
use_gpu: bool = False,
onnx_providers: Optional[List[str]] = None,
):
self.rec_image_shape = rec_image_shape
self.rec_batch_num = rec_batch_num
# ONNX session
if onnx_providers is None:
onnx_providers = (
["CUDAExecutionProvider", "CPUExecutionProvider"]
if use_gpu
else ["CPUExecutionProvider"]
)
sess_options = ort.SessionOptions()
sess_options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
self.session = ort.InferenceSession(
rec_onnx, sess_options=sess_options, providers=onnx_providers
)
self.input_name = self.session.get_inputs()[0].name
# CTC decoder
char_list = _load_char_dict(char_dict)
self._decoder = _CTCLabelDecode(char_list, use_space_char=True)
# ---- pre / post ----
def _preprocess(
self, img_list: List[np.ndarray]
) -> List[np.ndarray]:
"""Convert a list of crops into batch tensors (grouped by self.rec_batch_num)."""
num = len(img_list)
width_list = [im.shape[1] / float(im.shape[0]) for im in img_list]
indices = np.argsort(np.array(width_list))
batches = []
index_maps = [] # each element: list of original indices in this batch
for beg in range(0, num, self.rec_batch_num):
end = min(num, beg + self.rec_batch_num)
imgC, imgH, imgW = self.rec_image_shape
max_wh_ratio = imgW / imgH
for ino in range(beg, end):
orig_idx = indices[ino]
h, w = img_list[orig_idx].shape[:2]
max_wh_ratio = max(max_wh_ratio, w / h)
norm_list = []
idx_list = []
for ino in range(beg, end):
orig_idx = indices[ino]
norm = _resize_norm_img(
img_list[orig_idx],
self.rec_image_shape,
max_wh_ratio=max_wh_ratio,
)
norm_list.append(np.expand_dims(norm, axis=0))
idx_list.append(orig_idx)
if norm_list:
batches.append(np.concatenate(norm_list, axis=0).astype(np.float32))
index_maps.append(idx_list)
return batches, index_maps
def _postprocess(
self,
batch_outputs: List[np.ndarray],
index_maps: List[List[int]],
total_num: int,
) -> List[Tuple[str, float]]:
results = [("", 0.0)] * total_num
# Decode each batch separately — different batches may have different T.
for batch_preds, idx_list in zip(batch_outputs, index_maps):
texts = self._decoder(batch_preds)
for i, orig_idx in enumerate(idx_list):
results[orig_idx] = texts[i]
return results
# ---- public API ----
def __call__(
self, img: Union[np.ndarray, List[np.ndarray]]
) -> List[Tuple[str, float]]:
if isinstance(img, np.ndarray):
img = [img]
if not img:
return []
batches, index_maps = self._preprocess(img)
outputs = []
for batch in batches:
out = self.session.run(None, {self.input_name: batch})
outputs.append(out[0])
return self._postprocess(outputs, index_maps, len(img))
def predict_image(self, path: str) -> Tuple[str, float]:
im = cv2.imread(path)
if im is None:
raise FileNotFoundError(f"Cannot read: {path}")
return self.__call__(im)[0]
# ============================================================================
# 4. Evaluation
# ============================================================================
def evaluate(
ocr: PPOCRv6RecOnnx,
label_file: str,
dataset_root: str = "",
ignore_space: bool = True,
verbose: bool = False,
) -> dict:
"""Evaluate recognition accuracy against a ground-truth label file.
Label file format (one per line, tab-separated)::
rel/path/to/crop.png<TAB>ground truth text
The full image path is ``os.path.join(dataset_root, rel_path)``.
Args:
ocr: PPOCRv6RecOnnx instance.
label_file: path to tab-separated label file.
dataset_root: prefix directory for image paths in label file.
ignore_space: strip spaces before comparing.
verbose: print per-sample prediction details.
Returns:
dict with keys: ``acc``, ``norm_edit_dis``, ``total``, ``correct``,
``per_sample`` (list of per-sample details).
"""
images = []
targets = []
with open(label_file, "r", encoding="utf-8") as f:
for line in f:
line = line.strip()
if not line:
continue
parts = line.split("\t")
if len(parts) < 2:
continue
img_path = os.path.join(dataset_root, parts[0].strip())
images.append(img_path)
targets.append(parts[1].strip())
total = len(images)
if total == 0:
print("[WARN] No samples found in label file.")
return {"acc": 0.0, "norm_edit_dis": 0.0, "total": 0, "correct": 0, "per_sample": []}
# Load all images
imgs = []
for p in images:
im = cv2.imread(p)
if im is None:
print(f"[WARN] Cannot read {p}, skipping.")
imgs.append(np.zeros((32, 100, 3), dtype=np.uint8))
else:
imgs.append(im)
# Batch inference
rec_results = ocr(imgs)
correct = 0
total_edit_dis = 0.0
per_sample = []
for i, ((pred, conf), gt) in enumerate(zip(rec_results, targets)):
pred_clean = pred.replace(" ", "") if ignore_space else pred
gt_clean = gt.replace(" ", "") if ignore_space else gt
dist, norm_dist = _edit_distance(pred_clean, gt_clean)
total_edit_dis += norm_dist
is_correct = pred_clean == gt_clean
if is_correct:
correct += 1
sample = {
"image": images[i],
"pred": pred,
"gt": gt,
"confidence": round(conf, 4),
"correct": is_correct,
"edit_distance": int(dist),
"norm_edit_dis": round(norm_dist, 4),
}
per_sample.append(sample)
if verbose:
status = " OK" if is_correct else "MIS"
print(
f"[{status}] pred={pred!r:<30} gt={gt!r:<30} "
f"conf={conf:.4f} edit={int(dist)} ndis={norm_dist:.4f}"
)
acc = correct / total
norm_edit_dis = 1.0 - total_edit_dis / total
return {
"acc": round(acc, 6),
"norm_edit_dis": round(norm_edit_dis, 6),
"total": total,
"correct": correct,
"per_sample": per_sample,
}
# ============================================================================
# 5. CLI
# ============================================================================
def main():
parser = argparse.ArgumentParser(
description="PP-OCRv6 Recognition ONNX – inference & evaluation"
)
# Model
parser.add_argument(
"--rec_onnx", type=str,
default="onnx/rec_inference_static_sim.onnx",
help="Path to recognition ONNX model",
)
parser.add_argument(
"--char_dict", type=str,
default="onnx/rec_inference.yml",
help="Character dictionary: .yml (PostProcess.character_dict), .txt, or comma-list",
)
parser.add_argument("--batch_size", type=int, default=1, help="Recognition batch size")
parser.add_argument("--use_gpu", action="store_true", help="Enable GPU inference")
parser.add_argument("--rec_image_shape", type=str, default="3,48,320",
help="Recognition input shape C,H,W (comma separated)")
# Single image mode
parser.add_argument("--image", type=str, default=None, help="Single crop image path")
# Evaluation mode
parser.add_argument("--label_file", type=str,
default='dataset/ocr_rec_dataset_examples/val.txt',
help="Label file (image_path<TAB>gt_text per line)")
parser.add_argument("--dataset_root", type=str,
default="dataset/ocr_rec_dataset_examples",
help="Prefix directory for image paths in label file")
# Common
parser.add_argument("--ignore_space", action="store_true", default=True,
help="Ignore spaces when comparing (default: True)")
parser.add_argument("--verbose", action="store_true", help="Print per-sample results")
parser.add_argument("--output_json", type=str, default=None,
help="Save results to JSON file")
args = parser.parse_args()
# Build engine
char_dict_src = args.char_dict
if char_dict_src.startswith("[") or ("," in char_dict_src and not os.path.exists(char_dict_src)):
char_dict = [c.strip() for c in char_dict_src.split(",") if c.strip()]
else:
char_dict = char_dict_src
image_shape = tuple(int(v) for v in args.rec_image_shape.split(","))
if len(image_shape) != 3:
raise ValueError("--rec_image_shape requires 3 comma-separated integers")
ocr = PPOCRv6RecOnnx(
rec_onnx=args.rec_onnx,
char_dict=char_dict,
rec_image_shape=image_shape,
rec_batch_num=args.batch_size,
use_gpu=args.use_gpu,
)
# Single image mode
if args.image:
text, conf = ocr.predict_image(args.image)
print(f"text={text!r} confidence={conf:.4f}")
if args.output_json:
import json
with open(args.output_json, "w", encoding="utf-8") as f:
json.dump({"text": text, "confidence": conf}, f, ensure_ascii=False, indent=2)
return
# Evaluation mode
if args.label_file:
metrics = evaluate(
ocr,
args.label_file,
dataset_root=args.dataset_root,
ignore_space=args.ignore_space,
verbose=args.verbose,
)
print()
print("=" * 60)
print("Evaluation Results")
print("=" * 60)
print(f" Total samples: {metrics['total']}")
print(f" Correct (exact match): {metrics['correct']}")
print(f" Accuracy: {metrics['acc']:.4f} ({metrics['acc']*100:.2f}%)")
print(f" Norm Edit Distance: {metrics['norm_edit_dis']:.4f}")
print("=" * 60)
if args.output_json:
import json
out = {k: v for k, v in metrics.items() if k != "per_sample"}
out["per_sample"] = metrics["per_sample"]
with open(args.output_json, "w", encoding="utf-8") as f:
json.dump(out, f, ensure_ascii=False, indent=2)
print(f"\nResults saved to: {args.output_json}")
return
parser.error("Either --image or --label_file must be provided.")
if __name__ == "__main__":
main()