ethix's picture
feat: performance overhaul, client-side report, result cache, streaming cards, UX polish
05f720a
Raw History Blame Contribute Delete
5.04 kB
import os
import json
import uuid
import time
import threading
from datetime import datetime
from PIL import Image, ImageDraw, ImageFont
from huggingface_hub import HfApi, get_token
DATASET_REPO = "buildborderless/deepfake-explainability"
LOCAL_LOG_DIR = os.path.join(os.path.dirname(__file__), "logs")
def create_composite_xai_grid(image: Image.Image, results: list) -> Image.Image:
"""Combine original image and all XAI heatmap overlays into a single composite image grid."""
w, h = 384, 384
n = len(results) + 1 # Original + N methods
cols = min(n, 4)
rows = (n + cols - 1) // cols
grid_img = Image.new("RGB", (cols * w, rows * h), color=(15, 23, 42))
# Paste original
orig_resized = image.convert("RGB").resize((w, h))
grid_img.paste(orig_resized, (0, 0))
# Paste XAI overlays
for idx, res in enumerate(results, start=1):
r = idx // cols
c = idx % cols
overlay_b64 = res.overlay_b64 if hasattr(res, "overlay_b64") else res.get("overlay_b64")
if overlay_b64:
import base64
from io import BytesIO
b64_data = overlay_b64.split(",")[-1]
overlay_pil = Image.open(BytesIO(base64.b64decode(b64_data))).resize((w, h))
grid_img.paste(overlay_pil, (c * w, r * h))
return grid_img
def _async_push_to_dataset(record: dict, image: Image.Image, composite_grid: Image.Image):
"""Background worker function for dataset logging."""
os.makedirs(LOCAL_LOG_DIR, exist_ok=True)
log_file = os.path.join(LOCAL_LOG_DIR, "runs.jsonl")
# 1. Always append record to local JSONL
try:
with open(log_file, "a", encoding="utf-8") as f:
f.write(json.dumps(record) + "\n")
except Exception as e:
print(f"[logger] Local log write failed: {e}")
# 2. Save images locally
run_id = record["run_id"]
if image is not None:
img_path = os.path.join(LOCAL_LOG_DIR, f"{run_id}_input.png")
image.save(img_path)
grid_path = os.path.join(LOCAL_LOG_DIR, f"{run_id}_xai_composite.png")
composite_grid.save(grid_path)
# 3. Attempt Hub dataset push if token present
token = get_token() or os.environ.get("HF_TOKEN")
if token:
try:
api = HfApi(token=token)
# Upload local log files to dataset repo
api.upload_file(
path_or_fileobj=grid_path,
path_in_repo=f"grids/{run_id}_xai_composite.png",
repo_id=DATASET_REPO,
repo_type="dataset",
)
if image is not None:
api.upload_file(
path_or_fileobj=img_path,
path_in_repo=f"images/{run_id}_input.png",
repo_id=DATASET_REPO,
repo_type="dataset",
)
api.upload_file(
path_or_fileobj=log_file,
path_in_repo="runs.jsonl",
repo_id=DATASET_REPO,
repo_type="dataset",
)
print(f"[logger] Successfully pushed run {run_id} to dataset {DATASET_REPO}")
except Exception as e:
print(f"[logger] Hub push error (fallback to local): {e}")
def log_run_dataset(
image: Image.Image,
pred_data: dict,
results: list,
forensics_results: dict = None,
ground_truth: str = "unknown",
opt_in_image: bool = True,
) -> str:
"""
Log run details, prediction, ground truth, and XAI composite heatmaps.
Returns run_id.
"""
run_id = str(uuid.uuid4())[:8]
now_str = datetime.utcnow().isoformat()
is_correct = None
if ground_truth in ("real", "fake"):
pred_clean = pred_data.get("prediction", "").lower()
is_correct = (pred_clean == ground_truth)
methods_run = [r.name if hasattr(r, 'name') else r.get('name') for r in results]
per_method_time = {
(r.name if hasattr(r, 'name') else r.get('name')): (r.compute_time_ms if hasattr(r, 'compute_time_ms') else r.get('compute_time_ms'))
for r in results
}
record = {
"timestamp": now_str,
"run_id": run_id,
"prediction": pred_data.get("prediction"),
"probability": pred_data.get("probability"),
"confidence_pct": pred_data.get("confidence_pct"),
"ground_truth": ground_truth,
"is_correct": is_correct,
"opt_in_image": opt_in_image,
"methods_run": methods_run,
"per_method_time_ms": per_method_time,
"total_time_ms": sum(per_method_time.values()),
}
if forensics_results is not None:
record["forensics_results"] = forensics_results
composite_grid = create_composite_xai_grid(image, results)
logged_image = image if opt_in_image else None
# Spawn background thread to avoid blocking UI response
thread = threading.Thread(
target=_async_push_to_dataset,
args=(record, logged_image, composite_grid),
daemon=True
)
thread.start()
return run_id