ocr_annotator/app.py
2026-06-24 13:01:26 +04:00

373 lines
12 KiB
Python

import os
import csv
import io
import json
import zipfile
import threading
from pathlib import Path
from flask import Flask, render_template, request, jsonify, send_file, send_from_directory
app = Flask(__name__)
IMAGES_DIR = Path(os.environ.get("IMAGES_DIR", "/data/images"))
DATA_DIR = Path(os.environ.get("OUTPUT_FILE", "/data/labels.tsv")).parent
SUPPORTED_EXTENSIONS = {".jpg", ".jpeg", ".png", ".bmp", ".tiff", ".tif", ".webp"}
# TrOCR model used for automatic pre-labeling (Russian-finetuned by default).
TROCR_MODEL = os.environ.get("TROCR_MODEL", "raxtemur/trocr-base-ru")
_locks = {}
_locks_mutex = threading.Lock()
# Lazily-loaded TrOCR model, shared across requests in the worker process.
_model_cache = {}
_model_mutex = threading.Lock()
def get_lock(subfolder):
with _locks_mutex:
if subfolder not in _locks:
_locks[subfolder] = threading.Lock()
return _locks[subfolder]
def get_subfolders():
if not IMAGES_DIR.exists():
return []
subs = []
for d in sorted(IMAGES_DIR.iterdir()):
if d.is_dir():
has_images = any(
f.suffix.lower() in SUPPORTED_EXTENSIONS for f in d.iterdir() if f.is_file()
)
if has_images:
subs.append(d.name)
return subs
def get_images_in(subfolder):
folder = IMAGES_DIR / subfolder
if not folder.exists():
return []
return sorted([
f.name for f in folder.iterdir()
if f.is_file() and f.suffix.lower() in SUPPORTED_EXTENSIONS
])
def labels_path(subfolder):
return DATA_DIR / f"labels_{subfolder}.tsv"
def load_labels(subfolder):
"""Return {filename: {"text": str, "auto": bool}}.
TSV format: filename<TAB>text[<TAB>"auto"]. A third column "auto" marks
labels produced by the automatic pre-labeling pass (need a human recheck).
Rows with two columns are treated as human-verified for backward compat.
"""
path = labels_path(subfolder)
labels = {}
if path.exists():
with open(path, "r", encoding="utf-8") as f:
for row in csv.reader(f, delimiter="\t"):
if len(row) >= 2:
auto = len(row) >= 3 and row[2] == "auto"
labels[row[0]] = {"text": row[1], "auto": auto}
return labels
def save_labels(subfolder, labels):
path = labels_path(subfolder)
path.parent.mkdir(parents=True, exist_ok=True)
with open(path, "w", encoding="utf-8", newline="") as f:
writer = csv.writer(f, delimiter="\t")
for filename, entry in sorted(labels.items()):
row = [filename, entry["text"]]
if entry.get("auto"):
row.append("auto")
writer.writerow(row)
@app.route("/")
def index():
subfolders = get_subfolders()
return render_template("index.html", subfolders=subfolders)
@app.route("/api/subfolders")
def api_subfolders():
subs = get_subfolders()
result = []
for s in subs:
images = get_images_in(s)
labels = load_labels(s)
labeled = sum(1 for img in images if img in labels)
auto = sum(1 for img in images if labels.get(img, {}).get("auto"))
result.append({"name": s, "total": len(images), "labeled": labeled, "auto": auto})
return jsonify(result)
@app.route("/api/images/<subfolder>")
def api_images(subfolder):
images = get_images_in(subfolder)
labels = load_labels(subfolder)
return jsonify([
{
"filename": img,
"labeled": img in labels,
"text": labels.get(img, {}).get("text", ""),
"auto": labels.get(img, {}).get("auto", False),
}
for img in images
])
@app.route("/api/label/<subfolder>", methods=["POST"])
def api_label(subfolder):
data = request.json
filename = data.get("filename")
text = (data.get("text") or "").strip()
if not filename:
return jsonify({"error": "filename required"}), 400
images = get_images_in(subfolder)
if filename not in images:
return jsonify({"error": "image not found"}), 404
lock = get_lock(subfolder)
with lock:
labels = load_labels(subfolder)
if text == "" and filename in labels:
del labels[filename]
else:
# A manual save always counts as human-verified (clears the auto flag).
labels[filename] = {"text": text, "auto": False}
save_labels(subfolder, labels)
images_after = get_images_in(subfolder)
labels_after = load_labels(subfolder)
labeled = sum(1 for img in images_after if img in labels_after)
return jsonify({"ok": True, "labeled": labeled, "total": len(images_after)})
@app.route("/api/delete/<subfolder>", methods=["POST"])
def api_delete(subfolder):
data = request.json
filename = data.get("filename")
if not filename:
return jsonify({"error": "filename required"}), 400
# Validate: must be inside the subfolder, no path traversal
image_path = (IMAGES_DIR / subfolder / filename).resolve()
allowed_root = (IMAGES_DIR / subfolder).resolve()
if not str(image_path).startswith(str(allowed_root)):
return jsonify({"error": "invalid path"}), 400
if not image_path.exists():
return jsonify({"error": "file not found"}), 404
lock = get_lock(subfolder)
with lock:
# Remove from disk
image_path.unlink()
# Remove label if exists
labels = load_labels(subfolder)
if filename in labels:
del labels[filename]
save_labels(subfolder, labels)
images_after = get_images_in(subfolder)
labels_after = load_labels(subfolder)
labeled = sum(1 for img in images_after if img in labels_after)
return jsonify({"ok": True, "labeled": labeled, "total": len(images_after)})
@app.route("/images/<subfolder>/<path:filename>")
def serve_image(subfolder, filename):
return send_from_directory(IMAGES_DIR / subfolder, filename)
@app.route("/api/export")
def api_export():
"""Export a ZIP: labeled images + labels.tsv"""
subfolders = get_subfolders()
rows = []
image_paths = [] # (arc_name, fs_path)
for sub in subfolders:
labels = load_labels(sub)
for filename, entry in sorted(labels.items()):
text = entry["text"]
fs_path = IMAGES_DIR / sub / filename
if fs_path.exists():
arc_name = f"{sub}/{filename}"
rows.append((arc_name, text))
image_paths.append((arc_name, fs_path))
if not rows:
return jsonify({"error": "No labeled images yet"}), 404
buf = io.BytesIO()
with zipfile.ZipFile(buf, "w", zipfile.ZIP_DEFLATED) as zf:
# Write TSV
tsv_buf = io.StringIO()
writer = csv.writer(tsv_buf, delimiter="\t")
for row in rows:
writer.writerow(row)
zf.writestr("labels.tsv", tsv_buf.getvalue().encode("utf-8"))
# Write images
for arc_name, fs_path in image_paths:
zf.write(fs_path, arc_name)
buf.seek(0)
return send_file(
buf,
mimetype="application/zip",
as_attachment=True,
download_name="dataset.zip",
)
@app.route("/api/stats")
def api_stats():
subs = get_subfolders()
total = labeled = auto = 0
for s in subs:
imgs = get_images_in(s)
lbls = load_labels(s)
total += len(imgs)
labeled += sum(1 for img in imgs if img in lbls)
auto += sum(1 for img in imgs if lbls.get(img, {}).get("auto"))
return jsonify({"total": total, "labeled": labeled, "remaining": total - labeled, "auto": auto})
# ── Automatic pre-labeling with TrOCR ───────────────────────────────────────
def autolabel_status_path(subfolder):
return DATA_DIR / f"autolabel_{subfolder}.json"
def read_autolabel_status(subfolder):
"""Progress is persisted to a file so any gunicorn worker can read it."""
path = autolabel_status_path(subfolder)
if path.exists():
try:
return json.loads(path.read_text(encoding="utf-8"))
except Exception:
return None
return None
def write_autolabel_status(subfolder, status):
path = autolabel_status_path(subfolder)
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(json.dumps(status), encoding="utf-8")
def get_trocr():
"""Load (once) and return the TrOCR processor + model. Heavy imports are
done lazily so the app still runs in environments without torch installed."""
with _model_mutex:
if "model" not in _model_cache:
import torch
from transformers import TrOCRProcessor, VisionEncoderDecoderModel
processor = TrOCRProcessor.from_pretrained(TROCR_MODEL)
model = VisionEncoderDecoderModel.from_pretrained(TROCR_MODEL)
device = "cuda" if torch.cuda.is_available() else "cpu"
model.to(device)
model.eval()
_model_cache.update(
{"processor": processor, "model": model, "device": device, "torch": torch}
)
return (
_model_cache["processor"],
_model_cache["model"],
_model_cache["device"],
_model_cache["torch"],
)
def run_autolabel(subfolder):
"""Background job: run TrOCR over every unlabeled image in the subfolder and
store the predictions flagged as `auto` (i.e. needing a human recheck)."""
try:
images = get_images_in(subfolder)
existing = load_labels(subfolder)
todo = [img for img in images if img not in existing]
write_autolabel_status(
subfolder, {"state": "running", "total": len(todo), "done": 0, "error": None}
)
if not todo:
write_autolabel_status(
subfolder, {"state": "done", "total": 0, "done": 0, "error": None}
)
return
from PIL import Image
processor, model, device, torch = get_trocr()
lock = get_lock(subfolder)
done = 0
for img in todo:
try:
fs_path = IMAGES_DIR / subfolder / img
image = Image.open(fs_path).convert("RGB")
pixel_values = processor(images=image, return_tensors="pt").pixel_values.to(device)
with torch.no_grad():
generated_ids = model.generate(pixel_values, max_length=64)
text = processor.batch_decode(generated_ids, skip_special_tokens=True)[0].strip()
except Exception:
text = ""
# Save incrementally and never clobber a label a human added meanwhile.
with lock:
current = load_labels(subfolder)
if img not in current:
current[img] = {"text": text, "auto": True}
save_labels(subfolder, current)
done += 1
write_autolabel_status(
subfolder,
{"state": "running", "total": len(todo), "done": done, "error": None},
)
write_autolabel_status(
subfolder, {"state": "done", "total": len(todo), "done": done, "error": None}
)
except Exception as e:
write_autolabel_status(
subfolder, {"state": "error", "total": 0, "done": 0, "error": str(e)}
)
@app.route("/api/autolabel/<subfolder>", methods=["POST"])
def api_autolabel(subfolder):
if subfolder not in get_subfolders():
return jsonify({"error": "unknown subfolder"}), 404
status = read_autolabel_status(subfolder)
if status and status.get("state") == "running":
return jsonify({"error": "already running", "status": status}), 409
write_autolabel_status(
subfolder, {"state": "running", "total": 0, "done": 0, "error": None}
)
threading.Thread(target=run_autolabel, args=(subfolder,), daemon=True).start()
return jsonify({"ok": True})
@app.route("/api/autolabel/<subfolder>/status")
def api_autolabel_status(subfolder):
status = read_autolabel_status(subfolder) or {"state": "idle"}
return jsonify(status)
if __name__ == "__main__":
app.run(host="0.0.0.0", port=5000, debug=False)