Spaces:
Running on Zero
Running on Zero
File size: 4,451 Bytes
087643a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 | # Streamlit demo UI for Hugging Face Spaces. Reuses the same Detector and
# reasoning pipeline as the FastAPI app - no duplicated logic, just a
# visual layer for showing the thing actually working.
from __future__ import annotations
import os
from dotenv import load_dotenv
load_dotenv() # loads .env locally, no-op on HF Spaces (uses Space secrets instead)
import streamlit as st
from PIL import Image, ImageDraw
from app.constants import CLASS_NAMES, MODEL_VERSION
from app.detector import Detector
from app.reasoning.pipeline import answer_question
from scripts._render_utils import load_label_font
st.set_page_config(page_title="Document Layout Detection", page_icon="\U0001F4C4", layout="wide")
PALETTE = [
"#e6194b", "#3cb44b", "#ffe119", "#4363d8", "#f58231", "#911eb4",
"#46f0f0", "#f032e6", "#bcf60c", "#fabebe", "#008080",
]
@st.cache_resource
def get_detector() -> Detector:
# cache_resource so weights load once per container, not per request
detector = Detector()
try:
detector.load()
except FileNotFoundError:
pass # surfaced in the UI below instead of crashing the app
return detector
def draw_detections(image: Image.Image, detections) -> Image.Image:
annotated = image.copy()
draw = ImageDraw.Draw(annotated)
font = load_label_font(18)
for det in detections:
colour = PALETTE[det.class_id % len(PALETTE)]
box = det.bbox
draw.rectangle([box.x1, box.y1, box.x2, box.y2], outline=colour, width=3)
label = f"{det.class_name} {det.confidence:.2f}"
text_box = draw.textbbox((box.x1, box.y1), label, font=font)
draw.rectangle(
[text_box[0] - 2, text_box[1] - 2, text_box[2] + 2, text_box[3] + 2],
fill=colour,
)
draw.text((box.x1, box.y1), label, font=font, fill="white")
return annotated
detector = get_detector()
st.title("Constrained Document Layout Detection")
st.caption(f"RT-DETR fine-tuned on DocLayNet - {MODEL_VERSION}")
if not detector.is_loaded:
st.warning(
f"Model weights not found at `{os.environ.get('MODEL_PATH', './weights/best.pt')}`. "
"Detection and Q&A won't work until weights are available - see the README "
"for the download link, or set the MODEL_PATH secret on this Space."
)
with st.sidebar:
st.subheader("Classes")
st.write(", ".join(CLASS_NAMES))
st.subheader("Model")
st.write("loaded" if detector.is_loaded else "not loaded")
if not os.environ.get("GROQ_API_KEY"):
st.info("GROQ_API_KEY not set - the Ask tab needs it for the reasoning layer.")
tab_detect, tab_ask = st.tabs(["Detect", "Ask"])
with tab_detect:
st.write("Upload a document page to see the detected layout regions.")
uploaded = st.file_uploader("Image", type=["png", "jpg", "jpeg"], key="detect_upload")
if uploaded and st.button("Run detection", disabled=not detector.is_loaded):
image = Image.open(uploaded).convert("RGB")
with st.spinner("Running RT-DETR..."):
detections, inference_ms = detector.predict(image)
col1, col2 = st.columns(2)
col1.image(image, caption="Original", use_container_width=True)
col2.image(draw_detections(image, detections), caption="Detections", use_container_width=True)
st.caption(f"{len(detections)} detections in {inference_ms:.1f} ms")
if detections:
st.table([
{"class": d.class_name, "confidence": round(d.confidence, 3)}
for d in sorted(detections, key=lambda d: -d.confidence)
])
with tab_ask:
st.write("Ask a question about the document's layout - not its text content.")
uploaded_q = st.file_uploader("Image", type=["png", "jpg", "jpeg"], key="ask_upload")
question = st.text_input("Question", placeholder="How many tables are on this page?")
ask_disabled = not detector.is_loaded or not os.environ.get("GROQ_API_KEY")
if uploaded_q and question and st.button("Ask", disabled=ask_disabled):
image = Image.open(uploaded_q).convert("RGB")
with st.spinner("Thinking..."):
response = answer_question(image=image, question=question, detector=detector)
if response.insufficient_information:
st.warning(response.answer)
else:
st.success(response.answer)
with st.expander("Reasoning trace"):
st.json(response.reasoning_trace)
|