PP-DocLayoutV3-ONNX / pp_doclayout_v3_onnx.py
phungpx's picture
Upload ONNX export
409ef8b verified
Raw
History Blame Contribute Delete
14.7 kB
"""PP-DocLayoutV3 inference on ONNX Runtime — no torch, no transformers.
Pre- and post-processing are ported from
transformers/models/pp_doclayout_v3/image_processing_pp_doclayout_v3.py
so results match the PyTorch pipeline (boxes, labels, reading order, polygons).
from pp_doclayout_v3_onnx import PPDocLayoutV3ONNX
det = PPDocLayoutV3ONNX("pp_doclayoutv3.onnx", device="cuda")
for r in det.predict("page.jpg"):
print(r["order"], r["label"], r["score"], r["box"])
"""
from __future__ import annotations
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Sequence
import cv2
import numpy as np
import onnxruntime as ort
INPUT_SIZE = 800 # image processor resizes to a fixed 800x800 square
MASK_STRIDE = 4 # mask head resolution is 800/4 = 200
RESCALE_FACTOR = 1.0 / 255.0
# image_mean = [0,0,0], image_std = [1,1,1] -> normalisation is a no-op
ID2LABEL = {
0: "abstract", 1: "algorithm", 2: "aside_text", 3: "chart", 4: "content",
5: "formula", 6: "doc_title", 7: "figure_title", 8: "footer", 9: "footer",
10: "footnote", 11: "formula_number", 12: "header", 13: "header", 14: "image",
15: "formula", 16: "number", 17: "paragraph_title", 18: "reference",
19: "reference_content", 20: "seal", 21: "table", 22: "text", 23: "text",
24: "vision_footnote",
}
@dataclass
class Timings:
preprocess: float = 0.0
inference: float = 0.0
postprocess: float = 0.0
@property
def total(self) -> float:
return self.preprocess + self.inference + self.postprocess
def __str__(self) -> str:
return (
f"pre={self.preprocess * 1e3:.1f}ms infer={self.inference * 1e3:.1f}ms "
f"post={self.postprocess * 1e3:.1f}ms total={self.total * 1e3:.1f}ms"
)
# --------------------------------------------------------------------------- #
# Preprocessing
# --------------------------------------------------------------------------- #
def load_image_rgb(image: Any) -> np.ndarray:
"""Accept a path, a PIL image, or an HWC array. Returns RGB uint8."""
if isinstance(image, (str, Path)):
bgr = cv2.imread(str(image), cv2.IMREAD_COLOR)
if bgr is None:
raise FileNotFoundError(f"Could not read image: {image}")
return cv2.cvtColor(bgr, cv2.COLOR_BGR2RGB)
if isinstance(image, np.ndarray):
if image.ndim == 2:
return cv2.cvtColor(image, cv2.COLOR_GRAY2RGB)
if image.shape[2] == 4:
return cv2.cvtColor(image, cv2.COLOR_RGBA2RGB)
return image
return np.asarray(image.convert("RGB")) # PIL
def preprocess(images: Sequence[np.ndarray]) -> tuple[np.ndarray, list[tuple[int, int]]]:
"""Resize to 800x800 (bicubic, no antialias) and scale to [0, 1] NCHW float32.
The HF processor uses torchvision resize with antialias=False specifically to
approximate cv2.resize, so cv2 INTER_CUBIC is the reference behaviour here.
"""
batch = np.empty((len(images), 3, INPUT_SIZE, INPUT_SIZE), dtype=np.float32)
target_sizes: list[tuple[int, int]] = []
for i, img in enumerate(images):
h, w = img.shape[:2]
target_sizes.append((h, w))
resized = cv2.resize(img, (INPUT_SIZE, INPUT_SIZE), interpolation=cv2.INTER_CUBIC)
batch[i] = resized.astype(np.float32).transpose(2, 0, 1) * RESCALE_FACTOR
return batch, target_sizes
# --------------------------------------------------------------------------- #
# Post-processing (ported 1:1 from PPDocLayoutV3ImageProcessor)
# --------------------------------------------------------------------------- #
def _sigmoid(x: np.ndarray) -> np.ndarray:
return 1.0 / (1.0 + np.exp(-x, dtype=np.float64)).astype(np.float32)
def get_order_seqs(order_logits: np.ndarray) -> np.ndarray:
"""(B, Q, Q) pointer logits -> (B, Q) reading-order rank per query."""
scores = _sigmoid(order_logits)
batch_size, seq_len, _ = scores.shape
votes = np.triu(scores, 1).sum(axis=1) + np.tril(
1.0 - scores.transpose(0, 2, 1), -1
).sum(axis=1)
pointers = np.argsort(votes, axis=1, kind="stable")
order_seq = np.empty_like(pointers)
ranks = np.broadcast_to(np.arange(seq_len), (batch_size, seq_len))
np.put_along_axis(order_seq, pointers, ranks, axis=1)
return order_seq
def _extract_custom_vertices(polygon: np.ndarray, sharp_angle_thresh: float = 45) -> list[tuple]:
poly = np.array(polygon)
n = len(poly)
res = []
for i in range(n):
previous_point = poly[(i - 1) % n]
current_point = poly[i]
next_point = poly[(i + 1) % n]
v1 = previous_point - current_point
v2 = next_point - current_point
cross = (v1[1] * v2[0]) - (v1[0] * v2[1])
if cross < 0:
n1, n2 = np.linalg.norm(v1), np.linalg.norm(v2)
if n1 == 0 or n2 == 0:
res.append(tuple(current_point))
continue
angle = np.degrees(np.arccos(np.clip((v1 @ v2) / (n1 * n2), -1.0, 1.0)))
if abs(angle - sharp_angle_thresh) < 1:
direction = v1 / n1 + v2 / n2
direction = direction / np.linalg.norm(direction)
step = (n1 + n2) / 2
res.append(tuple(current_point + direction * step))
else:
res.append(tuple(current_point))
return res
def _mask2polygon(mask: np.ndarray, epsilon_ratio: float = 0.004):
contours, _ = cv2.findContours(mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
if not contours:
return None
contour = max(contours, key=cv2.contourArea)
epsilon = epsilon_ratio * cv2.arcLength(contour, True)
approx = cv2.approxPolyDP(contour, epsilon, True)
points = np.atleast_2d(approx.squeeze())
return _extract_custom_vertices(points)
def _extract_polygons(boxes: np.ndarray, masks: np.ndarray, scale_ratio) -> list:
scale_w, scale_h = scale_ratio[0] / MASK_STRIDE, scale_ratio[1] / MASK_STRIDE
mask_h, mask_w = masks.shape[1:]
polygons = []
for i in range(len(boxes)):
x_min, y_min, x_max, y_max = boxes[i].astype(np.int32)
box_w, box_h = int(x_max - x_min), int(y_max - y_min)
rect = np.array(
[[x_min, y_min], [x_max, y_min], [x_max, y_max], [x_min, y_max]], dtype=np.float32
)
if box_w <= 0 or box_h <= 0:
polygons.append(rect)
continue
x_start, x_end = np.clip(
[int(round(float(x_min * scale_w))), int(round(float(x_max * scale_w)))], 0, mask_w
)
y_start, y_end = np.clip(
[int(round(float(y_min * scale_h))), int(round(float(y_max * scale_h)))], 0, mask_h
)
cropped = masks[i, y_start:y_end, x_start:x_end]
if cropped.size == 0 or cropped.sum() == 0:
polygons.append(rect)
continue
resized = cv2.resize(cropped.astype(np.uint8), (box_w, box_h), interpolation=cv2.INTER_NEAREST)
polygon = _mask2polygon(resized)
if polygon is None or len(polygon) < 4:
polygons.append(rect)
continue
polygons.append(np.array(polygon, dtype=np.float32) + np.array([x_min, y_min]))
return polygons
def postprocess(
logits: np.ndarray,
pred_boxes: np.ndarray,
order_logits: np.ndarray,
out_masks: np.ndarray | None,
target_sizes: Sequence[tuple[int, int]],
threshold: float = 0.5,
) -> list[list[dict]]:
"""Returns one list of detections per image, already sorted by reading order."""
order_seqs = get_order_seqs(order_logits)
# cxcywh (normalised) -> xyxy (absolute)
centers, dims = pred_boxes[..., :2], pred_boxes[..., 2:]
boxes = np.concatenate([centers - 0.5 * dims, centers + 0.5 * dims], axis=-1)
sizes = np.asarray(target_sizes, dtype=np.float32) # (B, 2) as (h, w)
scale = np.stack([sizes[:, 1], sizes[:, 0], sizes[:, 1], sizes[:, 0]], axis=1)
boxes = boxes * scale[:, None, :]
batch_size, num_queries, num_classes = logits.shape
scores_all = _sigmoid(logits)
results: list[list[dict]] = []
for b in range(batch_size):
flat = scores_all[b].reshape(-1)
# torch.topk(k=num_queries) over the flattened (query, class) grid
top = np.argpartition(-flat, num_queries - 1)[:num_queries]
top = top[np.argsort(-flat[top], kind="stable")]
scores = flat[top]
labels = top % num_classes
query_idx = top // num_classes
keep = scores >= threshold
scores, labels, query_idx = scores[keep], labels[keep], query_idx[keep]
order = order_seqs[b][query_idx]
srt = np.argsort(order, kind="stable")
scores, labels, query_idx, order = scores[srt], labels[srt], query_idx[srt], order[srt]
sel_boxes = boxes[b][query_idx]
if out_masks is not None and len(sel_boxes):
masks = (_sigmoid(out_masks[b][query_idx]) > threshold).astype(np.uint8)
h, w = target_sizes[b]
polygons = _extract_polygons(sel_boxes, masks, [INPUT_SIZE / w, INPUT_SIZE / h])
else:
polygons = [
np.array([[x0, y0], [x1, y0], [x1, y1], [x0, y1]], dtype=np.float32)
for x0, y0, x1, y1 in sel_boxes
]
results.append(
[
{
"order": int(o),
"label_id": int(l),
"label": ID2LABEL.get(int(l), str(l)),
"score": float(s),
"box": [round(float(v), 2) for v in box],
"polygon": poly,
}
for o, l, s, box, poly in zip(order, labels, scores, sel_boxes, polygons)
]
)
return results
# --------------------------------------------------------------------------- #
# Engine
# --------------------------------------------------------------------------- #
class PPDocLayoutV3ONNX:
"""ONNX Runtime session, created once and reused."""
def __init__(
self,
onnx_path: str | Path,
*,
device: str = "cpu",
device_id: int = 0,
intra_op_num_threads: int | None = None,
threshold: float = 0.5,
trt_cache: str | None = None,
warmup: bool = True,
) -> None:
so = ort.SessionOptions()
so.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
if intra_op_num_threads:
so.intra_op_num_threads = intra_op_num_threads
providers: list[Any] = []
if device == "tensorrt":
providers.append((
"TensorrtExecutionProvider",
{
"device_id": device_id,
"trt_fp16_enable": True,
"trt_engine_cache_enable": bool(trt_cache),
"trt_engine_cache_path": trt_cache or "",
},
))
if device in ("cuda", "tensorrt"):
providers.append(("CUDAExecutionProvider", {"device_id": device_id}))
providers.append("CPUExecutionProvider")
self.session = ort.InferenceSession(str(onnx_path), sess_options=so, providers=providers)
self.input_name = self.session.get_inputs()[0].name
self.output_names = [o.name for o in self.session.get_outputs()]
self.has_masks = "out_masks" in self.output_names
self.threshold = threshold
self.last_timings = Timings()
if warmup:
self.session.run(
None, {self.input_name: np.zeros((1, 3, INPUT_SIZE, INPUT_SIZE), dtype=np.float32)}
)
@property
def providers(self) -> list[str]:
return self.session.get_providers()
def predict(
self, images: Any, threshold: float | None = None
) -> list[dict] | list[list[dict]]:
"""One image -> list of detections. A list of images -> list of those lists."""
single = not isinstance(images, (list, tuple))
image_list = [images] if single else list(images)
threshold = self.threshold if threshold is None else threshold
t0 = time.perf_counter()
rgb = [load_image_rgb(im) for im in image_list]
batch, target_sizes = preprocess(rgb)
t1 = time.perf_counter()
outputs = self.session.run(None, {self.input_name: batch})
t2 = time.perf_counter()
named = dict(zip(self.output_names, outputs))
results = postprocess(
named["logits"],
named["pred_boxes"],
named["order_logits"],
named.get("out_masks"),
target_sizes,
threshold=threshold,
)
t3 = time.perf_counter()
self.last_timings = Timings(t1 - t0, t2 - t1, t3 - t2)
return results[0] if single else results
def draw(image_path: str | Path, detections: list[dict], out_path: str | Path) -> None:
"""Quick visual sanity check: polygons + reading-order index."""
img = cv2.imread(str(image_path))
for det in detections:
poly = np.asarray(det["polygon"], dtype=np.int32).reshape(-1, 1, 2)
cv2.polylines(img, [poly], True, (0, 165, 255), 2)
x0, y0 = int(det["box"][0]), int(det["box"][1])
cv2.putText(
img, f"{det['order']}:{det['label']} {det['score']:.2f}",
(x0, max(y0 - 5, 12)), cv2.FONT_HERSHEY_SIMPLEX, 0.6, (0, 0, 255), 2,
)
cv2.imwrite(str(out_path), img)
if __name__ == "__main__":
import argparse
import json
p = argparse.ArgumentParser(description="PP-DocLayoutV3 ONNX Runtime inference")
p.add_argument("--onnx", required=True)
p.add_argument("--image", required=True, nargs="+")
p.add_argument("--device", default="cpu", choices=["cpu", "cuda", "tensorrt"])
p.add_argument("--threshold", type=float, default=0.5)
p.add_argument("--threads", type=int, default=None)
p.add_argument("--draw", default=None, help="write an annotated copy of the first image")
args = p.parse_args()
det = PPDocLayoutV3ONNX(
args.onnx, device=args.device, intra_op_num_threads=args.threads, threshold=args.threshold
)
print(f"providers: {det.providers}")
results = det.predict(args.image)
if not isinstance(results[0], list):
results = [results]
for path, dets in zip(args.image, results):
print(f"\n=== {path} === ({det.last_timings})")
for d in dets:
print(f" Order {d['order'] + 1}: {d['label']} {d['score']:.2f} {d['box']}")
if args.draw:
draw(args.image[0], results[0], args.draw)
print(f"\nannotated -> {args.draw}")
print(json.dumps({"count": [len(r) for r in results]}))