efwfe's picture
Upload folder using huggingface_hub
fe44a6e verified
Raw
History Blame Contribute Delete
5.77 kB
#!/usr/bin/env python3
"""
Example: PaddleOCR-VL Layer-12 Feature Extraction with ONNX
============================================================
Demonstrates:
1. Loading the ONNX model
2. Extracting features from an image
3. Computing quality via distance from reference
4. CV quality metrics (complementary)
5. Feature sensitivity to degradation
Requirements:
pip install onnxruntime numpy Pillow opencv-python
"""
from __future__ import annotations
import sys, os
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from inference.onnx_inference import Layer12ONNXExtractor
from inference.preprocessing import preprocess_for_onnx
import cv2
import numpy as np
from PIL import Image, ImageFilter
# ---------------------------------------------------------------------------
# 1. Load model
# ---------------------------------------------------------------------------
MODEL_PATH = os.path.join(
os.path.dirname(os.path.dirname(os.path.abspath(__file__))),
"model.onnx",
)
print("Loading ONNX model...")
extractor = Layer12ONNXExtractor(MODEL_PATH)
print(f" Feature dimension: {extractor.feature_dim}D")
print(f" Provider: {extractor.provider}")
# ---------------------------------------------------------------------------
# 2. Feature extraction
# ---------------------------------------------------------------------------
# Create a simple test image
img = Image.new("RGB", (512, 512), color=(240, 240, 240))
# Draw some "text-like" lines
from PIL import ImageDraw
draw = ImageDraw.Draw(img)
for y in range(20, 500, 30):
draw.rectangle([30, y, 480, y + 4], fill=(30, 30, 30))
print("\nExtracting features...")
features = extractor.extract(img)
print(f" Shape: {features.shape}")
print(f" Mean: {features.mean():.4f}")
print(f" Std: {features.std():.4f}")
print(f" Min: {features.min():.4f}")
print(f" Max: {features.max():.4f}")
# ---------------------------------------------------------------------------
# 3. Quality via distance from reference
# ---------------------------------------------------------------------------
# Pristine reference (same image)
pristine = img.copy()
# Degraded version
blurred = img.filter(ImageFilter.GaussianBlur(radius=5))
dist = extractor.distance_from_reference(blurred, pristine)
quality = extractor.quality_score(blurred, reference=pristine)
print(f"\nQuality assessment:")
print(f" Blurred vs Pristine:")
print(f" Cosine distance: {dist:.6f}")
print(f" Quality score: {quality:.4f}")
# Self-comparison
self_dist = extractor.distance_from_reference(pristine, pristine)
self_quality = extractor.quality_score(pristine, reference=pristine)
print(f" Pristine vs Pristine:")
print(f" Cosine distance: {self_dist:.6f}")
print(f" Quality score: {self_quality:.4f}")
# ---------------------------------------------------------------------------
# 4. Degradation sensitivity sweep
# ---------------------------------------------------------------------------
print("\nDegradation sensitivity (layer_12):")
print(f" {'Degradation':20s} {'Distance':>10s} {'Quality':>10s}")
print(f" {'-'*42}")
# Test different blur levels
for blur_r in [0, 1, 3, 5, 9, 13]:
degraded = img.filter(ImageFilter.GaussianBlur(radius=blur_r))
dist = extractor.distance_from_reference(degraded, pristine)
quality = extractor.quality_score(degraded, reference=pristine)
print(f" {'blur_'+str(blur_r):20s} {dist:>10.6f} {quality:>10.4f}")
# ---------------------------------------------------------------------------
# 5. Complementary CV quality metrics
# ---------------------------------------------------------------------------
def cv_quality_metrics(pil_img: Image.Image) -> dict:
"""Fast traditional CV metrics (complement deep features)."""
gray = cv2.cvtColor(np.array(pil_img.convert("RGB")), cv2.COLOR_RGB2GRAY)
# Laplacian variance (blur detector)
lap = cv2.Laplacian(gray, cv2.CV_64F).var()
# Brightness deviation from ideal (128)
brightness_dev = abs(gray.mean() - 128) / 128
# Edge density
edges = cv2.Canny(gray, 50, 150)
edge_density = edges.sum() / edges.size
# High-frequency energy (FFT)
fft = np.fft.fft2(gray.astype(np.float32))
fft_shift = np.fft.fftshift(fft)
mag = np.abs(fft_shift)
h, w = mag.shape
ch, cw = h // 2, w // 2
r = min(h, w) // 4
y, x = np.ogrid[-ch:h-ch, -cw:w-cw]
high_freq_mask = (x*x + y*y) > (r*r)
hf_energy = mag[high_freq_mask].sum() / (mag.sum() + 1e-12)
# Contrast (IQR)
p25, p75 = np.percentile(gray, [25, 75])
contrast_iqr = (p75 - p25) / 255
return {
"laplacian_var": float(lap),
"brightness_dev": float(brightness_dev),
"edge_density": float(edge_density),
"high_freq_energy": float(hf_energy),
"contrast_iqr": float(contrast_iqr),
}
print("\nCV quality metrics:")
for label, img_obj in [("pristine", pristine), ("blurred_r5", blurred)]:
cv = cv_quality_metrics(img_obj)
print(f" {label}:")
for k, v in cv.items():
print(f" {k}: {v:.4f}")
# ---------------------------------------------------------------------------
# 6. Preprocessing details
# ---------------------------------------------------------------------------
print("\nPreprocessing details:")
pixel_values, position_ids = preprocess_for_onnx(img)
print(f" pixel_values: shape={pixel_values.shape}, dtype={pixel_values.dtype}")
print(f" position_ids: shape={position_ids.shape}, dtype={position_ids.dtype}")
print(f" Num patches: {pixel_values.shape[1]}")
print(f" Patch size: 14×14×3")
print(f" Grid: sqrt({pixel_values.shape[1]}) ≈ {int(np.sqrt(pixel_values.shape[1]))}")
print("\n✓ All examples completed successfully!")