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