GazeCorrect / app.py
anonymous-IA's picture
Upload 41 files
0a40ab8 verified
Raw
History Blame Contribute Delete
14.3 kB
"""GazeCorrect: gaze -> Gaussian attention/noise -> text-guided regeneration."""
from __future__ import annotations
# ZeroGPU must be imported before torch. Safe fallback for local/regular Spaces.
try:
import spaces
except ImportError:
class _Spaces:
@staticmethod
def GPU(function): return function
spaces = _Spaces()
import os
from functools import lru_cache
from pathlib import Path
# Gradio 4.44 calls Starlette's template API using its pre-0.29 signature.
# Some current Space base images provide the newer signature instead.
import starlette.templating as _starlette_templating
_original_template_response = _starlette_templating.Jinja2Templates.TemplateResponse
def _compatible_template_response(self, *args, **kwargs):
if args and isinstance(args[0], str) and len(args) >= 2 and isinstance(args[1], dict):
template = self.get_template(args[0])
return _starlette_templating._TemplateResponse(
template,
args[1],
status_code=args[2] if len(args) > 2 else kwargs.get("status_code", 200),
headers=kwargs.get("headers"),
media_type=kwargs.get("media_type"),
background=kwargs.get("background"),
)
return _original_template_response(self, *args, **kwargs)
_starlette_templating.Jinja2Templates.TemplateResponse = _compatible_template_response
import gradio as gr
# Gradio 4.44 can encounter JSON-schema ``additionalProperties: false`` in
# newer Space dependencies. Avoid generating an invalid API schema at startup.
try:
import gradio_client.utils as _gradio_client_utils
_schema_to_python_type = _gradio_client_utils._json_schema_to_python_type
def _safe_schema_to_python_type(schema, defs=None):
if not isinstance(schema, dict):
return "Any"
if not isinstance(schema.get("additionalProperties"), dict):
schema = {key: value for key, value in schema.items() if key != "additionalProperties"}
return _schema_to_python_type(schema, defs)
_gradio_client_utils._json_schema_to_python_type = _safe_schema_to_python_type
except Exception:
pass
import numpy as np
import pandas as pd
import torch
from PIL import Image, ImageDraw, ImageFilter
MODEL_ID = "stanfordmimi/RoentGen-v2"
NO_DURATION = "— no duration column —"
def draw_points(image, points):
if image is None: return None
output = image.convert("RGB").copy(); draw = ImageDraw.Draw(output)
radius = max(6, min(output.size) // 70)
for i, (x, y, weight) in enumerate(points):
r = radius * (0.6 + 0.6 * weight)
draw.ellipse((x-r, y-r, x+r, y+r), outline=(255, 55, 45), width=3)
draw.text((x+r+2, y-r), str(i + 1), fill=(255, 55, 45))
return output
def load_image(file):
if file is None: return None, [], None, gr.update(value=None, visible=False)
try:
path = file if isinstance(file, str) else file.name
image = Image.open(path).convert("RGB")
# Keep the uploaded filename so a multi-image gaze export can be
# filtered by its ``id`` column when it is applied.
return image, [], Path(path).name, gr.update(value=image, visible=True)
except Exception as exc:
gr.Warning(f"Could not read image: {exc}")
return None, [], None, gr.update(value=None, visible=False)
def click_gaze(image, points, weight, event: gr.SelectData):
if image is None: return points, gr.update()
x, y = event.index
points = points + [(float(x), float(y), float(weight))]
return points, draw_points(image, points)
def prepare_csv(file):
"""Read a CSV and expose its columns for explicit user mapping."""
hidden = (None, gr.update(visible=False), gr.update(choices=[], value=None),
gr.update(choices=[], value=None), gr.update(choices=[], value=None),
gr.update(choices=[], value=NO_DURATION),
gr.update(visible=False))
if file is None:
return hidden
try:
path = file if isinstance(file, str) else file.name
frame = pd.read_csv(path, sep=None, engine="python")
if frame.empty:
raise ValueError("CSV contains no rows.")
columns = [str(column) for column in frame.columns]
lower = {column.lower().strip(): column for column in columns}
id_guess = lower.get("id") or lower.get("image_id") or lower.get("image") or lower.get("filename") or columns[0]
x_guess = lower.get("x") or lower.get("gaze_x") or lower.get("fix_x") or columns[0]
y_guess = lower.get("y") or lower.get("gaze_y") or lower.get("fix_y") or columns[min(1, len(columns) - 1)]
duration_guess = lower.get("duration") or lower.get("weight") or lower.get("fixation_duration") or NO_DURATION
return (frame.to_json(orient="split"), gr.update(visible=True),
gr.update(choices=columns, value=id_guess),
gr.update(choices=columns, value=x_guess),
gr.update(choices=columns, value=y_guess),
gr.update(choices=[NO_DURATION] + columns, value=duration_guess),
gr.update(visible=True))
except Exception as exc:
gr.Warning(f"Could not read CSV: {exc}")
return hidden
def apply_csv(frame_json, id_col, x_col, y_col, duration_col, image, image_name):
if image is None:
gr.Warning("Upload an image before applying gaze CSV data.")
return gr.update(), gr.update()
if not frame_json or not id_col or not x_col or not y_col:
gr.Warning("Select the ID, X, and Y columns first.")
return gr.update(), gr.update()
try:
frame = pd.read_json(frame_json, orient="split")
if id_col not in frame.columns:
raise ValueError(f'ID column "{id_col}" was not found in the CSV.')
if not image_name:
raise ValueError("The uploaded image name is unavailable for matching the CSV id column.")
image_path = Path(image_name)
accepted_ids = {image_path.name.casefold(), image_path.stem.casefold()}
gaze_ids = frame[id_col].astype(str).str.strip().str.casefold()
frame = frame[gaze_ids.isin(accepted_ids)]
if frame.empty:
raise ValueError(
f'No gaze rows matched image "{image_path.name}" in the "{id_col}" column.'
)
x, y = frame[x_col].astype(float).to_numpy(), frame[y_col].astype(float).to_numpy()
duration = (frame[duration_col].astype(float).to_numpy()
if duration_col and duration_col != NO_DURATION else np.ones(len(x)))
w, h = image.size
if len(x) and min(x) >= 0 and min(y) >= 0 and max(x) <= 1.05 and max(y) <= 1.05: x, y = x*w, y*h
duration = duration / max(float(duration.max()), 1e-8)
points = [(float(np.clip(a, 0, w-1)), float(np.clip(b, 0, h-1)), float(c)) for a,b,c in zip(x,y,duration)]
gr.Info(f'Loaded {len(points)} gaze points for "{Path(image_name).name}" from CSV.')
return points, draw_points(image, points)
except Exception as exc:
gr.Warning(f"Could not apply CSV: {exc}")
return gr.update(), gr.update()
def clear_gaze(image):
return [], gr.update(value=image) if image else gr.update()
def attention(image, points, sigma):
w, h = image.size; yy, xx = np.mgrid[:h, :w]
sigma = max(1, float(sigma))
heat = np.zeros((h, w), dtype=np.float32)
for x, y, weight in points:
heat += max(weight, .05) * np.exp(-((xx-x)**2 + (yy-y)**2)/(2*sigma**2))
return heat / (heat.max() + 1e-8)
def preview(image, heat):
color = np.zeros((*heat.shape, 3), np.uint8); color[..., 0] = (heat*255).astype(np.uint8)
return Image.blend(image.convert("RGB"), Image.fromarray(color), .5)
def noisy_image(image, heat, degree, feather, seed):
mask = Image.fromarray((heat*255).astype(np.uint8), "L")
if feather: mask = mask.filter(ImageFilter.GaussianBlur(float(feather)))
alpha = np.asarray(mask, np.float32)[..., None]/255
source = np.asarray(image.convert("RGB"), np.float32)
# Blend toward a new Gaussian-noise image instead of merely adding noise
# to the source. At degree=1, pixels at the attention peak are entirely
# noise, so no part of the original anatomy remains visible there.
noise = np.random.default_rng(seed).normal(127.5, 70.0, source.shape)
noise_amount = np.clip(alpha * float(degree), 0, 1)
noised = source * (1 - noise_amount) + noise * noise_amount
return Image.fromarray(np.clip(noised, 0, 255).astype(np.uint8)), mask
def device_dtype():
return ("cuda", torch.float16) if torch.cuda.is_available() else ("cpu", torch.float32)
@lru_cache(maxsize=2)
def pipeline(token):
# RoentGen-v2 is published as a text-to-image DiffusionPipeline. Using
# StableDiffusionImg2ImgPipeline bypasses its supported inference path and
# can yield non-radiographic results for different random seeds.
from diffusers import DiffusionPipeline
device, dtype = device_dtype()
return DiffusionPipeline.from_pretrained(MODEL_ID, torch_dtype=dtype, token=token).to(device)
@spaces.GPU
def generate(image, points, description, sigma, degree, feather, steps, seed,
oauth_token: gr.OAuthToken | None = None, progress=gr.Progress()):
if image is None: return None, None, None, "Upload an image first."
if not points: return None, None, None, "Add clicked or CSV gaze points first."
token = oauth_token.token if oauth_token is not None else (
os.environ.get("HF_TOKEN") or os.environ.get("HUGGINGFACEHUB_API_TOKEN")
)
if not token:
return None, None, None, "Sign in with Hugging Face first, then click Generate."
seed = None if seed < 0 else int(seed)
try:
progress(.1, desc="Creating gaze attention")
heat = attention(image, points, sigma); gaze = preview(image, heat)
progress(.25, desc="Adding attention-weighted Gaussian noise")
noised, mask = noisy_image(image, heat, degree, feather, seed)
progress(.4, desc="Generating chest X-ray with RoentGen-v2")
device, _ = device_dtype(); generator = None if seed is None else torch.Generator(device=device).manual_seed(seed)
finding = description.strip() or "Normal chest radiograph."
prompt = f"Chest radiograph. {finding}"
output = pipeline(token)(prompt=prompt, guidance_scale=3.5,
num_inference_steps=int(steps), generator=generator).images[0].convert("RGB")
corrected = Image.composite(output.resize(image.size), image.convert("RGB"), mask)
return gaze, noised, corrected, "Completed."
except Exception as exc:
return None, None, None, "Generation failed. Accept RoentGen access and set HF_TOKEN. Error: " + str(exc)
with gr.Blocks(title="GazeCorrect") as demo:
gr.Markdown("# GazeCorrect\nImage + gaze clicks/CSV + disease description → attention noise → corrected regenerated image. Use chest X-rays only; research use only.")
image_state, points_state, csv_state, image_name_state = gr.State(None), gr.State([]), gr.State(None), gr.State(None)
gr.LoginButton("Sign in with Hugging Face")
with gr.Row():
with gr.Column():
upload = gr.File(label="1. Upload chest X-ray", file_types=[".png", ".jpg", ".jpeg", ".webp", ".bmp"])
panel = gr.Image(label="2. Add gaze points by clicking", type="pil", visible=False, interactive=False)
with gr.Row():
weight = gr.Slider(.1, 1, value=1, step=.1, label="Next click weight")
clear_button = gr.Button("Clear gaze", variant="secondary")
with gr.Accordion("Import gaze CSV", open=False):
gr.Markdown("Upload a CSV and select its image ID, fixation X/Y, and optional duration/weight columns. Only rows whose selected ID matches the uploaded image filename (or filename without its extension) are imported.")
csv_file = gr.File(label="Choose CSV", file_types=[".csv", ".tsv", ".txt"])
with gr.Row(visible=False) as csv_mapping:
id_column = gr.Dropdown(label="Image ID column")
x_column = gr.Dropdown(label="X column")
y_column = gr.Dropdown(label="Y column")
duration_column = gr.Dropdown(label="Duration / weight (optional)")
apply_csv_button = gr.Button("Apply CSV gaze points", visible=False, variant="secondary")
description = gr.Textbox(label="3. Disease / radiology description", placeholder="Example: Right lower-lobe opacity. No pleural effusion.")
with gr.Accordion("Settings", open=False):
sigma = gr.Slider(5, 250, value=50, step=1, label="Fixation heatmap σ (pixels)")
degree = gr.Slider(0, 1, value=.35, step=.05, label="Gaussian noise degree")
feather = gr.Slider(0, 20, value=4, step=1, label="Mask feather")
steps = gr.Slider(10, 50, value=25, step=1, label="Diffusion steps")
seed = gr.Number(value=42, precision=0, label="Seed (-1 random)")
button = gr.Button("4. Generate", variant="primary")
status = gr.Textbox(label="Generation status", interactive=False, lines=3)
with gr.Column():
gaze_out = gr.Image(label="Gaze attention map")
noise_out = gr.Image(label="Attention-weighted Gaussian-noise image")
corrected_out = gr.Image(label="Corrected regenerated image")
upload.upload(load_image, upload, [image_state, points_state, image_name_state, panel])
panel.select(click_gaze, [image_state, points_state, weight], [points_state, panel])
clear_button.click(clear_gaze, image_state, [points_state, panel])
csv_file.upload(prepare_csv, csv_file, [csv_state, csv_mapping, id_column, x_column, y_column, duration_column, apply_csv_button])
apply_csv_button.click(apply_csv, [csv_state, id_column, x_column, y_column, duration_column, image_state, image_name_state], [points_state, panel])
button.click(generate, [image_state, points_state, description, sigma, degree, feather, steps, seed], [gaze_out, noise_out, corrected_out, status])
demo.queue().launch(show_error=True)