559 lines
18 KiB
Python
559 lines
18 KiB
Python
"""
|
||
OCR Internal Service
|
||
|
||
提供标准整图 OCR 识别服务。
|
||
"""
|
||
import io
|
||
import asyncio
|
||
from typing import List, Optional
|
||
from fastapi import FastAPI, File, UploadFile
|
||
from fastapi.responses import HTMLResponse
|
||
from pydantic import BaseModel
|
||
from PIL import Image
|
||
import numpy as np
|
||
from datetime import datetime
|
||
import logging
|
||
|
||
from config.config import settings
|
||
from engine.ocr_engine import OCRProcessor
|
||
|
||
# 配置日志
|
||
logging.basicConfig(
|
||
level=logging.DEBUG if settings.debug else logging.INFO,
|
||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s"
|
||
)
|
||
logger = logging.getLogger(__name__)
|
||
|
||
app = FastAPI(title="OCR Internal Service", version="2.0.0")
|
||
|
||
# 初始化OCR处理器
|
||
ocr = OCRProcessor(
|
||
use_det=settings.ocr_use_det,
|
||
use_cls=settings.ocr_use_cls,
|
||
use_rec=settings.ocr_use_rec,
|
||
model_dir=settings.ocr_model_dir,
|
||
text_score=settings.ocr_text_score,
|
||
max_side_len=settings.ocr_max_side_len,
|
||
det_limit_side_len=settings.ocr_det_limit_side_len,
|
||
det_box_thresh=settings.ocr_det_box_thresh,
|
||
det_thresh=settings.ocr_det_thresh,
|
||
rec_batch_num=settings.ocr_rec_batch_num,
|
||
device=settings.ocr_device,
|
||
device_id=settings.ocr_device_id,
|
||
intra_threads=settings.ocr_intra_threads,
|
||
inter_threads=settings.ocr_inter_threads,
|
||
enable_vertical=settings.ocr_enable_vertical,
|
||
)
|
||
|
||
|
||
# ─── 响应模型 ─────────────────────────────────────────────
|
||
|
||
class TextRegion(BaseModel):
|
||
confidence: float
|
||
text: str
|
||
text_region: List[List[int]] # [[x1,y1], [x2,y2], [x3,y3], [x4,y4]]
|
||
|
||
|
||
class OCRResponse(BaseModel):
|
||
code: int
|
||
message: str
|
||
data: Optional[List[TextRegion]] = None
|
||
timestamp: str
|
||
|
||
|
||
# ─── 辅助函数 ─────────────────────────────────────────────
|
||
|
||
def _build_response(word_results) -> OCRResponse:
|
||
"""将 word_results 构建为响应,兼容嵌套和扁平两种格式"""
|
||
ocr_result = []
|
||
|
||
if not word_results:
|
||
return OCRResponse(
|
||
code=0, message="success", data=ocr_result,
|
||
timestamp=datetime.now().isoformat(),
|
||
)
|
||
|
||
for line in word_results:
|
||
if isinstance(line, (tuple, list)) and len(line) >= 3 and isinstance(line[0], str):
|
||
# 扁平 item: 单条 (text, conf, coords),line 本身就是 item
|
||
text, confidence, coords = line[0], line[1], line[2]
|
||
if coords is not None:
|
||
ocr_result.append(TextRegion(
|
||
confidence=confidence, text=text, text_region=coords,
|
||
))
|
||
elif isinstance(line, (tuple, list)):
|
||
# 嵌套结构: line = ((item1), (item2), ...)
|
||
for item in line:
|
||
if isinstance(item, (tuple, list)) and len(item) >= 3:
|
||
text, confidence, coords = item[0], item[1], item[2]
|
||
if coords is not None:
|
||
ocr_result.append(TextRegion(
|
||
confidence=confidence, text=text, text_region=coords,
|
||
))
|
||
|
||
return OCRResponse(
|
||
code=0,
|
||
message="success",
|
||
data=ocr_result,
|
||
timestamp=datetime.now().isoformat(),
|
||
)
|
||
|
||
|
||
def _error_response(msg: str) -> OCRResponse:
|
||
return OCRResponse(
|
||
code=500,
|
||
message=msg,
|
||
timestamp=datetime.now().isoformat(),
|
||
)
|
||
|
||
|
||
# ─── API 端点 ─────────────────────────────────────────────
|
||
|
||
@app.get("/ocr_system/health")
|
||
async def health_check():
|
||
"""健康检查"""
|
||
return {
|
||
"status": "ok",
|
||
"service": "ocr",
|
||
"version": "2.0.0",
|
||
"device": ocr.device,
|
||
"timestamp": datetime.now().isoformat(),
|
||
}
|
||
|
||
|
||
@app.post("/ocr_system/ocr", response_model=OCRResponse)
|
||
async def ocr_recognition(file: UploadFile = File(...)):
|
||
"""上传图片文件并执行整图 OCR 识别。"""
|
||
try:
|
||
# 读取文件
|
||
image_bytes = await file.read()
|
||
image_file = io.BytesIO(image_bytes)
|
||
image = Image.open(image_file)
|
||
image_np = np.array(image)
|
||
|
||
logger.info(f"收到OCR请求: size={image_np.shape}")
|
||
|
||
# 预处理(原图直通:灰度化+CLAHE+锐化会损失边缘小字细节,已实测降低召回)
|
||
loop = asyncio.get_event_loop()
|
||
im_np = image_np
|
||
|
||
logger.info("使用标准OCR模式")
|
||
result = await loop.run_in_executor(None, ocr.process_with_vertical, im_np)
|
||
|
||
if result is None:
|
||
return _error_response("OCR处理失败")
|
||
|
||
return _build_response(result.word_results)
|
||
|
||
except Exception as e:
|
||
logger.error(f"OCR处理异常: {str(e)}", exc_info=True)
|
||
return _error_response(f"处理异常: {str(e)}")
|
||
|
||
|
||
# ─── 可视化页面 ───────────────────────────────────────────
|
||
|
||
@app.get("/ocr_system/viewer", response_class=HTMLResponse)
|
||
async def viewer():
|
||
"""OCR结果可视化页面"""
|
||
return HTMLResponse(content=VIEWER_HTML)
|
||
|
||
|
||
VIEWER_HTML = r"""
|
||
<!DOCTYPE html>
|
||
<html lang="zh-CN">
|
||
<head>
|
||
<meta charset="UTF-8">
|
||
<meta name="viewport" content="width=device-width, initial-scale=1.0">
|
||
<title>OCR 可视化工具</title>
|
||
<style>
|
||
* { margin: 0; padding: 0; box-sizing: border-box; }
|
||
body {
|
||
font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, sans-serif;
|
||
background: #1a1a2e;
|
||
color: #e0e0e0;
|
||
height: 100vh;
|
||
display: flex; flex-direction: column;
|
||
overflow: hidden;
|
||
}
|
||
.header {
|
||
background: #16213e;
|
||
padding: 12px 24px;
|
||
display: flex; align-items: center; gap: 16px;
|
||
border-bottom: 1px solid #0f3460;
|
||
flex-shrink: 0;
|
||
}
|
||
.header h1 { font-size: 18px; font-weight: 600; }
|
||
/* ── 三栏布局 ── */
|
||
.main-layout {
|
||
display: flex; flex: 1; min-height: 0;
|
||
}
|
||
|
||
/* 左栏:控件 */
|
||
.panel-left {
|
||
width: 280px; min-width: 280px;
|
||
background: #16213e;
|
||
padding: 16px;
|
||
display: flex; flex-direction: column; gap: 12px;
|
||
border-right: 1px solid #0f3460;
|
||
overflow-y: auto; flex-shrink: 0;
|
||
}
|
||
.section-title {
|
||
font-size: 12px; font-weight: 600; color: #a0a0b0;
|
||
text-transform: uppercase; letter-spacing: 1px;
|
||
padding-top: 4px;
|
||
}
|
||
.upload-zone {
|
||
border: 2px dashed #0f3460; border-radius: 10px;
|
||
padding: 24px 16px; text-align: center; cursor: pointer;
|
||
transition: all .2s;
|
||
}
|
||
.upload-zone:hover { border-color: #53d8fb; background: rgba(83,216,251,0.05); }
|
||
.upload-zone.dragover { border-color: #53d8fb; background: rgba(83,216,251,0.1); }
|
||
.upload-zone .icon { font-size: 32px; margin-bottom: 4px; }
|
||
.upload-zone .text { font-size: 13px; color: #a0a0b0; }
|
||
.upload-zone .hint { font-size: 11px; color: #606080; margin-top: 2px; }
|
||
.upload-zone input { display: none; }
|
||
|
||
.btn {
|
||
padding: 9px 14px; border: none; border-radius: 6px;
|
||
cursor: pointer; font-size: 13px; font-weight: 500;
|
||
transition: all .2s; width: 100%;
|
||
}
|
||
.btn-primary { background: #53d8fb; color: #1a1a2e; }
|
||
.btn-primary:hover { background: #7ce4ff; }
|
||
.btn-primary:disabled { background: #303050; color: #606080; cursor: not-allowed; }
|
||
|
||
.status {
|
||
font-size: 12px; padding: 8px 10px; border-radius: 6px;
|
||
background: #0f3460; word-break: break-all;
|
||
}
|
||
.status.success { background: #0a3d2e; color: #4ade80; }
|
||
.status.error { background: #3d0a0a; color: #f87171; }
|
||
.status.loading { background: #3d2e0a; color: #fbbf24; }
|
||
|
||
.info-grid {
|
||
display: grid; grid-template-columns: 1fr 1fr; gap: 6px;
|
||
font-size: 12px;
|
||
}
|
||
.info-item { background: #0f3460; padding: 6px 10px; border-radius: 4px; }
|
||
.info-item .label { color: #606080; font-size: 10px; }
|
||
.info-item .value { color: #e0e0e0; font-weight: 500; }
|
||
|
||
/* 中栏:OCR 画布区 */
|
||
.panel-center {
|
||
flex: 1; display: flex; flex-direction: column;
|
||
background: #0d1117; min-width: 0;
|
||
}
|
||
.canvas-pane {
|
||
flex: 1; display: flex; align-items: center; justify-content: center;
|
||
position: relative; overflow: hidden;
|
||
}
|
||
.canvas-pane .pane-label {
|
||
position: absolute; top: 6px; left: 8px;
|
||
font-size: 11px; padding: 2px 6px; border-radius: 3px;
|
||
background: rgba(0,0,0,0.6); color: #a0a0b0;
|
||
z-index: 2; pointer-events: none;
|
||
}
|
||
.canvas-pane canvas {
|
||
max-width: 100%; max-height: 100%;
|
||
object-fit: contain; border-radius: 2px;
|
||
}
|
||
|
||
/* 右栏:文本列表 */
|
||
.panel-right {
|
||
width: 280px; min-width: 280px;
|
||
background: #16213e; border-left: 1px solid #0f3460;
|
||
display: flex; flex-direction: column; flex-shrink: 0;
|
||
}
|
||
.panel-right .section-title { padding: 16px 16px 8px; }
|
||
#text-list {
|
||
flex: 1; overflow-y: auto; padding: 0 12px 12px;
|
||
font-size: 12px;
|
||
}
|
||
.text-item {
|
||
padding: 6px 8px; margin-bottom: 4px;
|
||
background: #0f3460; border-radius: 4px;
|
||
cursor: pointer; transition: all .12s;
|
||
display: flex; justify-content: space-between; align-items: center;
|
||
}
|
||
.text-item:hover { background: #1a3a70; }
|
||
.text-item.highlight { background: #1a4a80; border: 1px solid #53d8fb; }
|
||
.text-item .text { flex: 1; overflow: hidden; text-overflow: ellipsis; white-space: nowrap; }
|
||
.text-item .conf {
|
||
font-size: 10px; color: #53d8fb; margin-left: 8px; flex-shrink: 0;
|
||
}
|
||
|
||
.empty-state {
|
||
text-align: center; color: #606080;
|
||
position: absolute; top: 50%; left: 50%; transform: translate(-50%,-50%);
|
||
}
|
||
.empty-state .icon { font-size: 48px; margin-bottom: 8px; }
|
||
.empty-state p { font-size: 14px; }
|
||
|
||
@media (max-width: 1100px) {
|
||
.panel-right { width: 220px; min-width: 220px; }
|
||
.panel-left { width: 240px; min-width: 240px; }
|
||
}
|
||
</style>
|
||
</head>
|
||
<body>
|
||
<div class="header">
|
||
<h1>OCR 可视化工具</h1>
|
||
</div>
|
||
|
||
<div class="main-layout">
|
||
<!-- 左栏 -->
|
||
<div class="panel-left">
|
||
<div class="section-title">上传图片</div>
|
||
<div class="upload-zone" id="upload-zone">
|
||
<div class="icon">+</div>
|
||
<div class="text">点击或拖拽上传</div>
|
||
<div class="hint">JPG / PNG / BMP / TIFF</div>
|
||
<input type="file" id="file-input" accept="image/*">
|
||
</div>
|
||
|
||
<button class="btn btn-primary" id="btn-recognize" disabled>开始识别</button>
|
||
|
||
<div id="status-area"></div>
|
||
|
||
<div class="info-grid" id="info-grid" style="display:none;">
|
||
<div class="info-item"><div class="label">识别数</div><div class="value" id="info-count">-</div></div>
|
||
<div class="info-item"><div class="label">耗时</div><div class="value" id="info-elapsed">-</div></div>
|
||
<div class="info-item"><div class="label">尺寸</div><div class="value" id="info-size">-</div></div>
|
||
</div>
|
||
</div>
|
||
|
||
<!-- 中栏 -->
|
||
<div class="panel-center" id="panel-center">
|
||
<div class="canvas-pane" id="pane-text">
|
||
<div class="pane-label">文字文本框</div>
|
||
<div class="empty-state" id="empty-text">
|
||
<div class="icon">T</div>
|
||
<p>上传图片开始识别</p>
|
||
</div>
|
||
<canvas id="text-canvas" style="display:none;"></canvas>
|
||
</div>
|
||
</div>
|
||
|
||
<!-- 右栏 -->
|
||
<div class="panel-right">
|
||
<div class="section-title">识别文本</div>
|
||
<div id="text-list">
|
||
<div style="color:#606080;text-align:center;padding-top:40px;">暂无识别结果</div>
|
||
</div>
|
||
</div>
|
||
</div>
|
||
|
||
<script>
|
||
let ocrData = [];
|
||
let imageElement = null;
|
||
let isProcessing = false;
|
||
|
||
const uploadZone = document.getElementById('upload-zone');
|
||
const fileInput = document.getElementById('file-input');
|
||
const btnRecognize = document.getElementById('btn-recognize');
|
||
const textCanvas = document.getElementById('text-canvas');
|
||
const textList = document.getElementById('text-list');
|
||
const statusArea = document.getElementById('status-area');
|
||
const infoGrid = document.getElementById('info-grid');
|
||
const emptyText = document.getElementById('empty-text');
|
||
|
||
// ── 文件选择 ──
|
||
uploadZone.addEventListener('click', () => fileInput.click());
|
||
fileInput.addEventListener('change', (e) => handleFile(e.target.files[0]));
|
||
uploadZone.addEventListener('dragover', (e) => { e.preventDefault(); uploadZone.classList.add('dragover'); });
|
||
uploadZone.addEventListener('dragleave', () => uploadZone.classList.remove('dragover'));
|
||
uploadZone.addEventListener('drop', (e) => {
|
||
e.preventDefault(); uploadZone.classList.remove('dragover');
|
||
if (e.dataTransfer.files[0]) handleFile(e.dataTransfer.files[0]);
|
||
});
|
||
|
||
function handleFile(file) {
|
||
if (!file || !file.type.startsWith('image/')) {
|
||
showStatus('请选择图片文件', 'error'); return;
|
||
}
|
||
const reader = new FileReader();
|
||
reader.onload = (e) => {
|
||
imageElement = new Image();
|
||
imageElement.onload = () => {
|
||
ocrData = [];
|
||
textList.innerHTML = '<div style="color:#606080;text-align:center;padding-top:40px;">暂无识别结果</div>';
|
||
infoGrid.style.display = 'none';
|
||
drawBaseImages();
|
||
btnRecognize.disabled = false;
|
||
showStatus(`已加载: ${file.name} (${imageElement.width}x${imageElement.height})`, 'success');
|
||
};
|
||
imageElement.src = e.target.result;
|
||
};
|
||
reader.readAsDataURL(file);
|
||
}
|
||
|
||
function drawBaseImages() {
|
||
if (!imageElement) return;
|
||
drawCanvas(textCanvas, emptyText);
|
||
}
|
||
|
||
function drawCanvas(canvas, emptyEl) {
|
||
const pane = canvas.parentElement;
|
||
const maxW = pane.clientWidth * 0.92;
|
||
const maxH = pane.clientHeight * 0.88;
|
||
let w = imageElement.width, h = imageElement.height;
|
||
const scale = Math.min(maxW / w, maxH / h, 1);
|
||
w = Math.floor(w * scale); h = Math.floor(h * scale);
|
||
|
||
canvas.width = w; canvas.height = h;
|
||
const ctx = canvas.getContext('2d');
|
||
ctx.drawImage(imageElement, 0, 0, w, h);
|
||
|
||
canvas.style.display = 'block';
|
||
if (emptyEl) emptyEl.style.display = 'none';
|
||
}
|
||
|
||
// ── OCR ──
|
||
btnRecognize.addEventListener('click', async () => {
|
||
if (!imageElement || isProcessing) return;
|
||
isProcessing = true;
|
||
btnRecognize.disabled = true;
|
||
btnRecognize.textContent = '识别中...';
|
||
showStatus('正在识别...', 'loading');
|
||
ocrData = [];
|
||
textList.innerHTML = '<div style="color:#606080;text-align:center;padding-top:40px;">识别中...</div>';
|
||
infoGrid.style.display = 'none';
|
||
|
||
try {
|
||
const blob = await fetch(imageElement.src).then(r => r.blob());
|
||
const formData = new FormData();
|
||
formData.append('file', blob, 'image.png');
|
||
|
||
const t0 = performance.now();
|
||
const resp = await fetch('/ocr_system/ocr', { method: 'POST', body: formData });
|
||
const t1 = performance.now();
|
||
const json = await resp.json();
|
||
|
||
if (json.code !== 0) {
|
||
showStatus(`识别失败: ${json.message}`, 'error'); return;
|
||
}
|
||
|
||
ocrData = json.data || [];
|
||
const elapsed = (t1 - t0).toFixed(0);
|
||
|
||
document.getElementById('info-count').textContent = ocrData.length;
|
||
document.getElementById('info-elapsed').textContent = `${elapsed}ms`;
|
||
document.getElementById('info-size').textContent = `${imageElement.width}x${imageElement.height}`;
|
||
infoGrid.style.display = 'grid';
|
||
|
||
redrawTextCanvas();
|
||
renderTextList();
|
||
|
||
if (ocrData.length > 0) {
|
||
showStatus(`完成: ${ocrData.length}条 / ${elapsed}ms`, 'success');
|
||
} else {
|
||
showStatus('未识别到文字', 'error');
|
||
}
|
||
} catch (err) {
|
||
showStatus(`请求失败: ${err.message}`, 'error');
|
||
} finally {
|
||
isProcessing = false;
|
||
btnRecognize.disabled = false;
|
||
btnRecognize.textContent = '开始识别';
|
||
}
|
||
});
|
||
|
||
// ── 文字框画布 ──
|
||
function redrawTextCanvas() {
|
||
if (!imageElement) return;
|
||
drawCanvas(textCanvas, emptyText);
|
||
|
||
if (ocrData.length === 0) return;
|
||
|
||
const ctx = textCanvas.getContext('2d');
|
||
const scaleX = textCanvas.width / imageElement.width;
|
||
const scaleY = textCanvas.height / imageElement.height;
|
||
const lw = Math.max(1, 1.5 * (textCanvas.width / 800));
|
||
|
||
ocrData.forEach((item, i) => {
|
||
const coords = item.text_region;
|
||
if (!coords || coords.length < 4) return;
|
||
ctx.strokeStyle = `hsl(${(i * 37) % 360}, 70%, 65%)`;
|
||
ctx.lineWidth = lw;
|
||
ctx.beginPath();
|
||
ctx.moveTo(coords[0][0] * scaleX, coords[0][1] * scaleY);
|
||
for (let j = 1; j < coords.length; j++) {
|
||
ctx.lineTo(coords[j][0] * scaleX, coords[j][1] * scaleY);
|
||
}
|
||
ctx.closePath();
|
||
ctx.stroke();
|
||
});
|
||
}
|
||
|
||
// ── 文本列表 ──
|
||
function renderTextList() {
|
||
if (ocrData.length === 0) {
|
||
textList.innerHTML = '<div style="color:#606080;text-align:center;padding-top:40px;">未识别到文字</div>';
|
||
return;
|
||
}
|
||
textList.innerHTML = ocrData.map((item, i) => `
|
||
<div class="text-item" data-index="${i}" onclick="highlightBox(${i})">
|
||
<span class="text">${escapeHtml(item.text)}</span>
|
||
<span class="conf">${(item.confidence * 100).toFixed(1)}%</span>
|
||
</div>
|
||
`).join('');
|
||
}
|
||
|
||
let highlightedIndex = -1;
|
||
function highlightBox(index) {
|
||
redrawTextCanvas();
|
||
const ctx = textCanvas.getContext('2d');
|
||
const scaleX = textCanvas.width / imageElement.width;
|
||
const scaleY = textCanvas.height / imageElement.height;
|
||
const coords = ocrData[index].text_region;
|
||
|
||
ctx.strokeStyle = '#ff6b6b';
|
||
ctx.lineWidth = 3;
|
||
ctx.beginPath();
|
||
ctx.moveTo(coords[0][0] * scaleX, coords[0][1] * scaleY);
|
||
for (let j = 1; j < coords.length; j++) {
|
||
ctx.lineTo(coords[j][0] * scaleX, coords[j][1] * scaleY);
|
||
}
|
||
ctx.closePath();
|
||
ctx.stroke();
|
||
|
||
document.querySelectorAll('.text-item').forEach(el => el.classList.remove('highlight'));
|
||
const target = document.querySelector(`.text-item[data-index="${index}"]`);
|
||
if (target) target.classList.add('highlight');
|
||
target?.scrollIntoView({ behavior: 'smooth', block: 'nearest' });
|
||
}
|
||
|
||
function escapeHtml(text) {
|
||
const d = document.createElement('div');
|
||
d.textContent = text;
|
||
return d.innerHTML;
|
||
}
|
||
|
||
function showStatus(msg, type) {
|
||
statusArea.innerHTML = `<div class="status ${type}">${msg}</div>`;
|
||
}
|
||
|
||
window.addEventListener('resize', () => {
|
||
if (imageElement && !isProcessing) {
|
||
redrawTextCanvas();
|
||
}
|
||
});
|
||
</script>
|
||
</body>
|
||
</html>
|
||
"""
|
||
|
||
|
||
# ─── 启动入口 ─────────────────────────────────────────────
|
||
|
||
if __name__ == "__main__":
|
||
import uvicorn
|
||
logger.info(f"启动OCR服务: {settings.host}:{settings.port}")
|
||
uvicorn.run(
|
||
"main:app",
|
||
host=settings.host,
|
||
port=settings.port,
|
||
reload=settings.debug
|
||
)
|