AutoLineDigitizer / api_server.py
t29mato's picture
Initial deploy: AutoLineDigitizer Gradio API for HuggingFace Spaces
27c91ef
Raw
History Blame
16.5 kB
# -*- coding: utf-8 -*-
"""
AutoLineDigitizer API Server for Hugging Face Spaces.
Exposes chart line extraction as a Gradio API.
"""
import sys
import os
# Setup paths (same as desktop_app.py)
SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__))
CHARTDETE_DIR = os.path.join(SCRIPT_DIR, "submodules", "chartdete")
LINEFORMER_DIR = os.path.join(SCRIPT_DIR, "submodules", "lineformer")
MMDET_DIR = os.path.join(LINEFORMER_DIR, "mmdetection")
SRC_DIR = os.path.join(SCRIPT_DIR, "src")
sys.path.insert(0, SCRIPT_DIR)
sys.path.insert(0, SRC_DIR)
sys.path.insert(0, MMDET_DIR)
sys.path.insert(0, LINEFORMER_DIR)
sys.path.insert(0, CHARTDETE_DIR)
# Register ChartDete custom models
CHARTDETE_AVAILABLE = False
try:
import mmdet # noqa: F401
from mmdet.models.roi_heads.cascade_roi_head_LGF import CascadeRoIHead_LGF # noqa: F401
CHARTDETE_AVAILABLE = True
except Exception as e:
print(f"ChartDete custom models not available: {e}")
import cv2
import numpy as np
import json
import io
import zipfile
import urllib.request
import ssl
import gradio as gr
from datetime import datetime, timezone
# ============================================================
# Model download & loading
# ============================================================
GITHUB_REPO = "t29mato/AutoLineDigitizer"
GITHUB_RELEASE_TAG = "models"
# Model files: filename -> download source
MODEL_FILES = {
"iter_3000.pth": {
"source": "github",
"url": f"https://github.com/{GITHUB_REPO}/releases/download/{GITHUB_RELEASE_TAG}/iter_3000.pth",
},
"checkpoint.pth": {
"source": "github",
"url": f"https://github.com/{GITHUB_REPO}/releases/download/{GITHUB_RELEASE_TAG}/checkpoint.pth",
},
}
_infer_module = None
_chartdete_module = None
def download_file(url, dest_path):
"""Download a file with progress logging and SSL fallback."""
print(f" Downloading {os.path.basename(dest_path)} from {url} ...")
tmp_path = dest_path + ".tmp"
req = urllib.request.Request(url)
try:
response_ctx = urllib.request.urlopen(req)
except Exception:
ctx = ssl._create_unverified_context()
response_ctx = urllib.request.urlopen(req, context=ctx)
with response_ctx as response:
total_size = int(response.headers.get("Content-Length", 0))
downloaded = 0
block_size = 1024 * 1024 # 1MB
with open(tmp_path, "wb") as f:
while True:
chunk = response.read(block_size)
if not chunk:
break
f.write(chunk)
downloaded += len(chunk)
if total_size > 0:
pct = downloaded * 100 // total_size
print(f" {os.path.basename(dest_path)}: {downloaded // (1024*1024)}MB / {total_size // (1024*1024)}MB ({pct}%)")
os.replace(tmp_path, dest_path)
print(f" {os.path.basename(dest_path)} downloaded successfully.")
def ensure_models():
"""Download model files if they don't exist."""
models_dir = os.path.join(SCRIPT_DIR, "models")
os.makedirs(models_dir, exist_ok=True)
for filename, info in MODEL_FILES.items():
dest = os.path.join(models_dir, filename)
if not os.path.exists(dest):
download_file(info["url"], dest)
else:
print(f" {filename} already exists, skipping download.")
def load_models():
"""Download (if needed) and load LineFormer and ChartDete models."""
global _infer_module, _chartdete_module
print("Checking model files...")
ensure_models()
models_dir = os.path.join(SCRIPT_DIR, "models")
# Load LineFormer
import infer
config_path = os.path.join(LINEFORMER_DIR, "lineformer_swin_t_config.py")
ckpt_path = os.path.join(models_dir, "iter_3000.pth")
infer.load_model(config_path, ckpt_path, "cpu")
_infer_module = infer
print("LineFormer model loaded.")
# Load ChartDete
if CHARTDETE_AVAILABLE:
import chartdete_infer
chartdete_config = os.path.join(SCRIPT_DIR, "config", "chartdete_config.py")
chartdete_ckpt = os.path.join(models_dir, "checkpoint.pth")
chartdete_infer.load_chartdete_model(
config_path=chartdete_config,
checkpoint_path=chartdete_ckpt,
device="cpu",
)
_chartdete_module = chartdete_infer
print("ChartDete model loaded.")
else:
print("ChartDete not available, axis detection disabled.")
def extract_lines(img):
"""Run LineFormer inference and return raw centerline points."""
line_dataseries = _infer_module.get_dataseries(img, to_clean=False)
raw_lines = []
for line in line_dataseries:
if len(line) == 0:
continue
raw_lines.append([[int(pt["x"]), int(pt["y"])] for pt in line])
return raw_lines
def arc_length_resample(points, n_points):
"""Resample at equidistant intervals along pixel-space arc length."""
pts = np.array(points, dtype=float)
diffs = np.diff(pts, axis=0)
seg_lengths = np.sqrt((diffs ** 2).sum(axis=1))
cum_arc = np.zeros(len(pts))
cum_arc[1:] = np.cumsum(seg_lengths)
total_length = cum_arc[-1]
if total_length == 0:
return [points[0]]
target_distances = np.linspace(0, total_length, n_points)
result = []
seg_idx = 0
for d in target_distances:
while seg_idx < len(seg_lengths) - 1 and cum_arc[seg_idx + 1] < d:
seg_idx += 1
seg_span = cum_arc[seg_idx + 1] - cum_arc[seg_idx]
t = 0.0 if seg_span == 0 else (d - cum_arc[seg_idx]) / seg_span
x = pts[seg_idx, 0] + t * (pts[seg_idx + 1, 0] - pts[seg_idx, 0])
y = pts[seg_idx, 1] + t * (pts[seg_idx + 1, 1] - pts[seg_idx, 1])
result.append([int(round(x)), int(round(y))])
return result
def downsample_points(points, mode="max_points", max_points=20, fixed_step=10):
"""Downsample points based on mode."""
if len(points) <= 1:
return points
if mode == "none":
return points
elif mode == "fixed":
return points[::fixed_step]
elif mode == "max_points":
if len(points) <= max_points:
return points
step = max(1, len(points) // max_points)
return points[::step]
elif mode == "arc_length":
if len(points) <= max_points:
return points
return arc_length_resample(points, max_points)
return points
def detect_axis_calibration(img):
"""Detect axis calibration using ChartDete + OCR."""
if _chartdete_module is None:
return None
detections = _chartdete_module.detect_chart_elements(img, score_thr=0.3)
axis_info = _chartdete_module.get_axis_info(detections, img=img, with_ocr=True)
calibration = axis_info.get("calibration")
if calibration is None:
return None
has_x = "x1_pixel" in calibration and "x2_pixel" in calibration
has_y = "y1_pixel" in calibration and "y2_pixel" in calibration
if not (has_x and has_y):
return None
plot_area = axis_info.get("plot_area")
if plot_area:
x_calib_y = plot_area[3]
y_calib_x = plot_area[0]
else:
x_calib_y = img.shape[0] * 0.9
y_calib_x = img.shape[1] * 0.1
return {
"x1_px": calibration["x1_pixel"],
"x1_py": x_calib_y,
"x1_val": calibration["x1_value"],
"x2_px": calibration["x2_pixel"],
"x2_py": x_calib_y,
"x2_val": calibration["x2_value"],
"y1_px": y_calib_x,
"y1_py": calibration["y2_pixel"],
"y1_val": calibration["y2_value"],
"y2_px": y_calib_x,
"y2_py": calibration["y1_pixel"],
"y2_val": calibration["y1_value"],
"xIsLogScale": False,
"yIsLogScale": False,
}
def convert_to_starry_digitizer_format(data_series, img_shape, axis_config=None):
"""Convert extracted data to StarryDigitizer project.json format."""
timestamp = datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%S.%f")[:-3] + "Z"
if axis_config is None:
axis_set = {
"id": 1, "name": "XY Axes 1",
"x1": {"name": "x1", "value": 0, "coord": {"xPx": 0, "yPx": float(img_shape[0])}},
"x2": {"name": "x2", "value": 100, "coord": {"xPx": float(img_shape[1]), "yPx": float(img_shape[0])}},
"y1": {"name": "y1", "value": 0, "coord": {"xPx": 0, "yPx": float(img_shape[0])}},
"y2": {"name": "y2", "value": 100, "coord": {"xPx": 0, "yPx": 0}},
"xIsLogScale": False, "yIsLogScale": False,
"considerGraphTilt": False, "pointMode": 0, "isVisible": True,
}
else:
axis_set = {
"id": 1, "name": "XY Axes 1",
"x1": {"name": "x1", "value": axis_config["x1_val"], "coord": {"xPx": axis_config["x1_px"], "yPx": axis_config["x1_py"]}},
"x2": {"name": "x2", "value": axis_config["x2_val"], "coord": {"xPx": axis_config["x2_px"], "yPx": axis_config["x2_py"]}},
"y1": {"name": "y1", "value": axis_config["y1_val"], "coord": {"xPx": axis_config["y1_px"], "yPx": axis_config["y1_py"]}},
"y2": {"name": "y2", "value": axis_config["y2_val"], "coord": {"xPx": axis_config["y2_px"], "yPx": axis_config["y2_py"]}},
"xIsLogScale": axis_config.get("xIsLogScale", False),
"yIsLogScale": axis_config.get("yIsLogScale", False),
"considerGraphTilt": False, "pointMode": 0, "isVisible": True,
}
datasets = [{
"id": 1, "name": "dataset 1", "axisSetId": 1,
"points": [], "visiblePointIds": [], "manuallyAddedPointIds": [],
}]
for idx, series in enumerate(data_series):
points = []
visible_ids = []
for pt_idx, pt in enumerate(series["points"]):
pt_id = pt_idx + 1
points.append({"id": pt_id, "xPx": float(pt[0]), "yPx": float(pt[1])})
visible_ids.append(pt_id)
datasets.append({
"id": idx + 2, "name": f"Line {idx + 1}", "axisSetId": 1,
"points": points, "visiblePointIds": visible_ids, "manuallyAddedPointIds": [],
})
return {
"version": "1.11.2", "timestamp": timestamp,
"axisSets": [axis_set], "activeAxisSetId": 1,
"datasets": datasets, "activeDatasetId": len(datasets),
"canvasHandler": {"scale": 1.0, "manualMode": 0},
}
# ============================================================
# Gradio API endpoint
# ============================================================
def digitize_chart(
image,
auto_axis_detection: bool = True,
downsample_mode: str = "arc_length",
max_points: int = 20,
fixed_step: int = 10,
sort_mode: str = "mean_y_desc",
output_format: str = "starry_digitizer_json",
):
"""
Extract line data from a chart image.
Args:
image: Input chart image (PIL Image from Gradio)
auto_axis_detection: Enable ChartDete + OCR axis detection
downsample_mode: "none", "max_points", "fixed", or "arc_length"
max_points: Max points per line (for max_points/arc_length modes)
fixed_step: Step size (for fixed mode)
sort_mode: "original", "mean_y_desc", or "mean_y_asc"
output_format: "starry_digitizer_json", "starry_digitizer_zip", or "json"
Returns:
For starry_digitizer_json/json: JSON string
For starry_digitizer_zip: ZIP file path
"""
if image is None:
return json.dumps({"error": "No image provided"})
# Convert PIL Image to BGR numpy array
img_rgb = np.array(image)
img_bgr = cv2.cvtColor(img_rgb, cv2.COLOR_RGB2BGR)
# Extract lines
raw_lines = extract_lines(img_bgr)
# Downsample
data_series = []
for all_points in raw_lines:
points = downsample_points(all_points, downsample_mode, max_points, fixed_step)
data_series.append({"points": points})
# Sort
if sort_mode == "mean_y_desc" and len(data_series) > 0:
data_series = sorted(data_series, key=lambda s: np.mean([pt[1] for pt in s["points"]]))
elif sort_mode == "mean_y_asc" and len(data_series) > 0:
data_series = sorted(data_series, key=lambda s: np.mean([pt[1] for pt in s["points"]]), reverse=True)
# Axis detection
axis_config = None
if auto_axis_detection and _chartdete_module is not None:
axis_config = detect_axis_calibration(img_bgr)
# Build output
if output_format == "starry_digitizer_zip":
project_json = convert_to_starry_digitizer_format(data_series, img_bgr.shape, axis_config)
zip_buffer = io.BytesIO()
with zipfile.ZipFile(zip_buffer, "w", zipfile.ZIP_DEFLATED) as zf:
_, img_encoded = cv2.imencode(".png", img_bgr)
zf.writestr("image.png", img_encoded.tobytes())
zf.writestr("project.json", json.dumps(project_json, indent=2, ensure_ascii=False))
zip_buffer.seek(0)
import tempfile
tmp = tempfile.NamedTemporaryFile(suffix=".zip", delete=False)
tmp.write(zip_buffer.read())
tmp.close()
return tmp.name
elif output_format == "starry_digitizer_json":
project_json = convert_to_starry_digitizer_format(data_series, img_bgr.shape, axis_config)
return json.dumps(project_json, ensure_ascii=False)
else: # "json" - raw extraction result
result = {
"num_lines": len(data_series),
"lines": [
{"line_index": i, "num_points": len(s["points"]), "points": s["points"]}
for i, s in enumerate(data_series)
],
"axis_config": axis_config,
"image_shape": {"height": img_bgr.shape[0], "width": img_bgr.shape[1]},
}
return json.dumps(result, ensure_ascii=False)
# ============================================================
# Gradio Interface
# ============================================================
def create_app():
"""Create Gradio app."""
with gr.Blocks(title="AutoLineDigitizer API") as demo:
gr.Markdown("# AutoLineDigitizer API")
gr.Markdown(
"Upload a chart image to automatically extract line data. "
"Results can be imported into [StarryDigitizer](https://digitizer.starrydata.org/)."
)
with gr.Row():
with gr.Column():
image_input = gr.Image(type="pil", label="Chart Image")
auto_axis = gr.Checkbox(value=True, label="Auto Axis Detection (ChartDete + OCR)")
downsample = gr.Dropdown(
choices=["none", "max_points", "fixed", "arc_length"],
value="arc_length",
label="Downsample Mode",
)
max_pts = gr.Slider(minimum=5, maximum=200, value=20, step=1, label="Max Points per Line")
fixed_stp = gr.Slider(minimum=1, maximum=50, value=10, step=1, label="Fixed Step")
sort = gr.Dropdown(
choices=["original", "mean_y_desc", "mean_y_asc"],
value="mean_y_desc",
label="Sort Mode",
)
out_fmt = gr.Dropdown(
choices=["starry_digitizer_json", "starry_digitizer_zip", "json"],
value="starry_digitizer_json",
label="Output Format",
)
run_btn = gr.Button("Extract Lines", variant="primary")
with gr.Column():
output = gr.Textbox(label="Result (JSON)", lines=20, max_lines=50)
file_output = gr.File(label="Download ZIP", visible=False)
def on_run(image, auto_axis, downsample, max_pts, fixed_stp, sort, out_fmt):
result = digitize_chart(image, auto_axis, downsample, max_pts, fixed_stp, sort, out_fmt)
if out_fmt == "starry_digitizer_zip":
return gr.update(value="ZIP file generated. Download below."), gr.update(value=result, visible=True)
else:
return gr.update(value=result), gr.update(visible=False)
run_btn.click(
fn=on_run,
inputs=[image_input, auto_axis, downsample, max_pts, fixed_stp, sort, out_fmt],
outputs=[output, file_output],
)
return demo
if __name__ == "__main__":
print("Loading models...")
load_models()
print("All models loaded. Starting Gradio server...")
demo = create_app()
demo.queue()
demo.launch(server_name="0.0.0.0", server_port=7860)