trocr-meter-ocr / app.py
sergsof2's picture
Update app.py
9fc2955 verified
Raw History Blame Contribute Delete
5.33 kB
#1 ячейка
from ultralytics import YOLO
from glob import glob
import os
import gradio as gr
import torch
from transformers import TrOCRProcessor, VisionEncoderDecoderModel
import cv2
import numpy as np
from PIL import Image
import re
#2 ячейка
# Пути к весам, если модель дообучена
# checkpoint = "/content/drive/MyDrive/T/TrOCR/trocr-small-finetuned"
checkpoint = "microsoft/trocr-small-printed" # если используешь оригинальную base small
#processor = TrOCRProcessor.from_pretrained(checkpoint)
processor = TrOCRProcessor.from_pretrained(checkpoint, use_fast=False)
model = VisionEncoderDecoderModel.from_pretrained(checkpoint)
model.to('cuda' if torch.cuda.is_available() else 'cpu')
model.eval()
def ocr_image(img, processor, model, device='cuda'):
# img — PIL.Image (или numpy, но надо привести к PIL!)
if isinstance(img, np.ndarray):
img = Image.fromarray(cv2.cvtColor(img, cv2.COLOR_BGR2RGB))
pixel_values = processor(img, return_tensors="pt").pixel_values.to(device)
with torch.no_grad():
generated_ids = model.generate(pixel_values)
return processor.batch_decode(generated_ids, skip_special_tokens=True)[0]
# 1. Загружаем модель и делаем инференс
model_yolo = YOLO('y_obb.pt')
def rotate_image_simple(img, angle):
h, w = img.shape[:2]
center = (w // 2, h // 2)
M = cv2.getRotationMatrix2D(center, angle, 1.0)
rotated = cv2.warpAffine(img, M, (w, h), flags=cv2.INTER_LINEAR, borderValue=(255,255,255))
return rotated, M
def transform_bbox_points(points, M):
# points: (4, 2) — четыре точки рамки
ones = np.ones((points.shape[0], 1))
points_hom = np.hstack([points, ones])
points_rot = M @ points_hom.T
return points_rot.T # (4, 2)
def filter_predicted_text(text):
# Оставить только цифры, точку, дробную черту и тире
allowed = r"[^0-9\.\-\/]" # . — точка, - — тире, / — дробная черта
f_text = re.sub(allowed, " ", text)
return f'{f_text}'
#3 ячейка
# ---- Функция для градио ----
def predict_number(upload_img):
# upload_img — PIL.Image
# Преобразуем PIL в numpy для пайплайна
img = np.array(upload_img)
img = cv2.cvtColor(img, cv2.COLOR_RGB2BGR)
# --- вставь сюда свой пайплайн обработки crop/roi! ---
results = model_yolo(img, verbose=False)
obbs = results[0].obb.data.cpu().numpy()
best_det = None
best_score = -1
for det in obbs:
x, y, w, h, theta, conf, cls = det
if int(cls) == 2 and conf > best_score:
best_det = det
best_score = conf
if best_det is not None:
x, y, w, h, theta, conf, cls = best_det
theta_deg = np.degrees(theta)
rect = ((x, y), (w, h), theta_deg)
box = cv2.boxPoints(rect)
# Оценить ориентацию на исходном изображении
w_box = np.linalg.norm(box[0] - box[1])
h_box = np.linalg.norm(box[1] - box[2])
if w_box <= h_box:
angle_to_rotate = theta_deg
else:
angle_to_rotate = theta_deg - 90
#print(f"Поворачиваем оригинал на {angle_to_rotate:.2f} градусов")
img_rot, M = rotate_image_simple(img, angle=angle_to_rotate)
# Перевести bbox на повернутую картинку
box_rot = transform_bbox_points(box, M)
box_rot_int = np.int32(box_rot)
# ----- Расширяем рамку на 10% -----
# 1. AABB
x_min, y_min = box_rot.min(axis=0)
x_max, y_max = box_rot.max(axis=0)
w_aabb = x_max - x_min
h_aabb = y_max - y_min
# 2. Добавим по 10% по бокам
pad_x = 0.05 * w_aabb
pad_y = 0.1 * h_aabb
x_min_pad = int(max(x_min - pad_x, 0))
y_min_pad = int(max(y_min - pad_y, 0))
x_max_pad = int(min(x_max + pad_x, img_rot.shape[1]))
y_max_pad = int(min(y_max + pad_y, img_rot.shape[0]))
# 3. Вырезаем ROI
roi = img_rot[y_min_pad:y_max_pad, x_min_pad:x_max_pad]
# OCR-предсказание (roi — numpy)
pred_text = ocr_image(roi, processor, model, device=('cuda' if torch.cuda.is_available() else 'cpu'))
filtered_text = filter_predicted_text(pred_text)
return filtered_text
else:
return "Не найден ни один объект с серийным номером"
#4 ячейка
with gr.Blocks() as demo:
gr.Markdown("## TrOCR: Распознавание номера на изображении")
with gr.Row():
with gr.Column():
img_input = gr.Image(type="pil", label="Загрузите изображение")
btn = gr.Button("Получить номер")
clear_btn = gr.Button("Очистить")
with gr.Column():
output_text = gr.Textbox(label="Результат")
btn.click(predict_number, inputs=img_input, outputs=output_text)
clear_btn.click(lambda: (None, ""), None, [img_input, output_text])
demo.launch()