#!/usr/bin/env python
"""
=============================================================================
AUTOMATIC LENS CORRECTION â€” COMPLETE KAGGLE NOTEBOOK
=============================================================================
Self-contained pipeline that runs all 4 phases on Kaggle's servers where
competition data is pre-mounted. No downloads needed.

Phases:
  1. Extract empirical warp map from training pairs (optical flow)
  2. Apply warp map to all test images (baseline)
  3. Per-image test-time optimization (maximize line straightness)
  4. Adaptive ensemble blending

Output: /kaggle/working/submission_ensemble.zip (upload to bounty.autohdr.com)

Estimated runtime: 60-90 minutes on Kaggle CPU (4 cores, 16GB RAM)
                   30-45 minutes on Kaggle GPU instance
=============================================================================
"""

import os
import sys
import time
import json
import zipfile
import random
import warnings
from pathlib import Path
from concurrent.futures import ProcessPoolExecutor, ThreadPoolExecutor, as_completed

import numpy as np
import cv2
from scipy.optimize import minimize

warnings.filterwarnings('ignore')

# ============================================================================
# CONFIGURATION
# ============================================================================

# Kaggle paths (competition data is auto-mounted here)
# Auto-discover data directories
KAGGLE_INPUT = Path("/kaggle/input")
COMPETITION_DIR = None
TRAIN_DIR = None
TEST_DIR = None

# Find competition data - could be at various paths
possible_roots = [
    KAGGLE_INPUT / "automatic-lens-correction",
    KAGGLE_INPUT / "competitions" / "automatic-lens-correction",
    KAGGLE_INPUT,
]

def find_data_dirs():
    """Auto-discover train and test directories."""
    global COMPETITION_DIR, TRAIN_DIR, TEST_DIR
    
    print("Auto-discovering data directories...")
    
    # List everything under /kaggle/input (3 levels deep)
    if KAGGLE_INPUT.exists():
        print(f"Contents of {KAGGLE_INPUT}:")
        for item in sorted(KAGGLE_INPUT.iterdir()):
            print(f"  {'[DIR]' if item.is_dir() else '[FILE]'} {item.name}")
            if item.is_dir():
                for sub in sorted(item.iterdir())[:15]:
                    print(f"    {'[DIR]' if sub.is_dir() else '[FILE]'} {sub.name}")
                    if sub.is_dir():
                        for subsub in sorted(sub.iterdir())[:5]:
                            print(f"      {'[DIR]' if subsub.is_dir() else '[FILE]'} {subsub.name}")
                        remaining = len(list(sub.iterdir())) - 5
                        if remaining > 0:
                            print(f"      ... and {remaining} more")
    
    # Search for train directory (contains *_original.jpg files)
    for root in possible_roots:
        if not root.exists():
            continue
        # Check root directly
        if list(root.glob("*_original.jpg"))[:1]:
            TRAIN_DIR = root
            COMPETITION_DIR = root
            break
        # Check one level down
        for subdir in sorted(root.iterdir()):
            if subdir.is_dir():
                if list(subdir.glob("*_original.jpg"))[:1]:
                    TRAIN_DIR = subdir
                    COMPETITION_DIR = root
                    break
                # Check two levels down
                for subsubdir in sorted(subdir.iterdir()):
                    if subsubdir.is_dir() and list(subsubdir.glob("*_original.jpg"))[:1]:
                        TRAIN_DIR = subsubdir
                        COMPETITION_DIR = subdir
                        break
                if TRAIN_DIR:
                    break
        if TRAIN_DIR:
            break
    
    # Search for test directory (any dir with .jpg but NOT _original.jpg)
    search_roots = [COMPETITION_DIR] if COMPETITION_DIR else possible_roots
    for root in search_roots:
        if root is None or not root.exists():
            continue
        for subdir in sorted(root.iterdir()):
            if subdir.is_dir() and subdir != TRAIN_DIR:
                jpgs = list(subdir.glob("*.jpg"))[:1]
                originals = list(subdir.glob("*_original.jpg"))[:1]
                if jpgs and not originals:
                    TEST_DIR = subdir
                    break
        if TEST_DIR:
            break
    
    print(f"\nDiscovered paths:")
    print(f"  COMPETITION_DIR: {COMPETITION_DIR}")
    print(f"  TRAIN_DIR: {TRAIN_DIR}")
    print(f"  TEST_DIR: {TEST_DIR}")

find_data_dirs()

# Output paths
WORKING_DIR = Path("/kaggle/working")
WARP_DIR = WORKING_DIR / "warp_maps"
BASELINE_DIR = WORKING_DIR / "outputs" / "baseline"
OPTIMIZED_DIR = WORKING_DIR / "outputs" / "optimized"
ENSEMBLE_DIR = WORKING_DIR / "outputs" / "ensemble"

for d in [WARP_DIR, BASELINE_DIR, OPTIMIZED_DIR, ENSEMBLE_DIR]:
    d.mkdir(parents=True, exist_ok=True)

# Phase 1 config
WORK_H = 512            # Working resolution height
WORK_W = 768            # Working resolution width
MAX_TRAIN_PAIRS = 2000  # Number of pairs for warp map (more = better but slower)
MAX_FLOW_MAG = 50.0     # Outlier rejection threshold
MIN_FLOW_MAG = 0.5

# Phase 3 config
OPTIM_RESOLUTION = 800  # Long edge for optimizer
MAX_ITERATIONS = 60     # Nelder-Mead iterations per image
MIN_LINES_FOR_OPTIM = 5
MIN_LINE_LENGTH = 40
K1_BOUNDS = (-0.5, 0.5)
K2_BOUNDS = (-0.3, 0.3)
CX_BOUNDS = (-0.05, 0.05)
CY_BOUNDS = (-0.05, 0.05)

# Phase 4 config
MIN_LINES_FULL_OPTIM = 30
MIN_LINES_ANY_OPTIM = 5
OPTIM_WEIGHT_BASE = 0.3
OPTIM_WEIGHT_MAX = 0.85

JPEG_QUALITY = 95
NUM_WORKERS = max(1, os.cpu_count() - 1)

print(f"System: {os.cpu_count()} CPUs, using {NUM_WORKERS} workers")
print(f"Train dir: {TRAIN_DIR} (exists: {TRAIN_DIR.exists() if TRAIN_DIR else 'N/A'})")
print(f"Test dir: {TEST_DIR} (exists: {TEST_DIR.exists() if TEST_DIR else 'N/A'})")

# ============================================================================
# PHASE 0: DATA DISCOVERY
# ============================================================================

def discover_data():
    """Find all training pairs and test images."""
    print("\n" + "=" * 60)
    print("PHASE 0: DATA DISCOVERY")
    print("=" * 60)
    
    # Find training pairs
    train_pairs = []
    if TRAIN_DIR and TRAIN_DIR.exists():
        train_files = sorted(TRAIN_DIR.glob("*_original.jpg"))
        for orig in train_files:
            gen = Path(str(orig).replace("_original.jpg", "_generated.jpg"))
            if gen.exists():
                train_pairs.append((str(orig), str(gen)))
    
    # Find test images
    test_images = []
    if TEST_DIR and TEST_DIR.exists():
        test_images = sorted(TEST_DIR.glob("*.jpg"))
    
    print(f"Training pairs: {len(train_pairs)}")
    print(f"Test images: {len(test_images)}")
    
    if train_pairs:
        # Look at a sample image to understand resolution
        sample = cv2.imread(str(train_pairs[0][0]))
        if sample is not None:
            print(f"Sample training image size: {sample.shape[1]}x{sample.shape[0]}")
    
    if test_images:
        sample = cv2.imread(str(test_images[0]))
        if sample is not None:
            print(f"Sample test image size: {sample.shape[1]}x{sample.shape[0]}")
    
    return train_pairs, test_images

# ============================================================================
# PHASE 1: EXTRACT EMPIRICAL WARP MAP
# ============================================================================

def compute_optical_flow(img_dist, img_corr):
    """Compute dense optical flow from distorted to corrected."""
    gray_d = cv2.cvtColor(img_dist, cv2.COLOR_BGR2GRAY)
    gray_c = cv2.cvtColor(img_corr, cv2.COLOR_BGR2GRAY)
    
    flow = cv2.calcOpticalFlowFarneback(
        gray_d, gray_c, None,
        pyr_scale=0.5, levels=5, winsize=15,
        iterations=5, poly_n=7, poly_sigma=1.5,
        flags=cv2.OPTFLOW_FARNEBACK_GAUSSIAN
    )
    return flow


def extract_warp_map(train_pairs):
    """Build averaged warp map from training pairs."""
    print("\n" + "=" * 60)
    print("PHASE 1: EXTRACTING EMPIRICAL WARP MAP")
    print("=" * 60)
    
    start_time = time.time()
    
    # Sample if too many
    pairs = train_pairs
    if len(pairs) > MAX_TRAIN_PAIRS:
        random.seed(42)
        pairs = random.sample(pairs, MAX_TRAIN_PAIRS)
        print(f"Sampled {MAX_TRAIN_PAIRS} pairs from {len(train_pairs)}")
    
    # Accumulate flow maps
    flow_sum = np.zeros((WORK_H, WORK_W, 2), dtype=np.float64)
    count = 0
    skipped = 0
    
    for i, (orig_path, gen_path) in enumerate(pairs):
        try:
            img_d = cv2.imread(orig_path)
            img_c = cv2.imread(gen_path)
            
            if img_d is None or img_c is None:
                skipped += 1
                continue
            
            # Resize to working resolution
            img_d = cv2.resize(img_d, (WORK_W, WORK_H), interpolation=cv2.INTER_AREA)
            img_c = cv2.resize(img_c, (WORK_W, WORK_H), interpolation=cv2.INTER_AREA)
            
            # Compute flow
            flow = compute_optical_flow(img_d, img_c)
            
            # Validate
            mag = np.sqrt(flow[:,:,0]**2 + flow[:,:,1]**2)
            mean_mag = np.mean(mag)
            
            if mean_mag > MAX_FLOW_MAG or mean_mag < MIN_FLOW_MAG:
                skipped += 1
                continue
            
            flow_sum += flow
            count += 1
            
            if (i + 1) % 200 == 0:
                elapsed = time.time() - start_time
                rate = (i + 1) / elapsed
                remaining = (len(pairs) - i - 1) / rate / 60
                print(f"  [{i+1}/{len(pairs)}] {count} valid, {skipped} skipped | "
                      f"{rate:.1f} pairs/sec | ~{remaining:.1f} min remaining")
        
        except Exception as e:
            skipped += 1
            continue
    
    if count == 0:
        print("ERROR: No valid pairs! Cannot build warp map.")
        return None
    
    mean_flow = (flow_sum / count).astype(np.float32)
    
    # Stats
    mag = np.sqrt(mean_flow[:,:,0]**2 + mean_flow[:,:,1]**2)
    print(f"\n  Warp map stats:")
    print(f"    Valid pairs: {count}, Skipped: {skipped}")
    print(f"    Mean displacement: {mag.mean():.3f} px")
    print(f"    Max displacement: {mag.max():.3f} px")
    print(f"    Time: {(time.time()-start_time)/60:.1f} minutes")
    
    # Save
    np.save(str(WARP_DIR / "mean_flow.npy"), mean_flow)
    
    # Save visualization
    hsv = np.zeros((*mean_flow.shape[:2], 3), dtype=np.uint8)
    mag_norm = cv2.normalize(mag, None, 0, 255, cv2.NORM_MINMAX).astype(np.uint8)
    _, ang = cv2.cartToPolar(mean_flow[:,:,0], mean_flow[:,:,1])
    hsv[:,:,0] = (ang * 180 / np.pi / 2).astype(np.uint8)
    hsv[:,:,1] = 255
    hsv[:,:,2] = mag_norm
    vis = cv2.cvtColor(hsv, cv2.COLOR_HSV2BGR)
    cv2.imwrite(str(WARP_DIR / "flow_vis.png"), vis)
    cv2.imwrite(str(WARP_DIR / "magnitude.png"), 
                cv2.applyColorMap(mag_norm, cv2.COLORMAP_JET))
    
    print(f"  Warp map saved to {WARP_DIR}")
    return mean_flow

# ============================================================================
# PHASE 2: BASELINE (APPLY WARP MAP)
# ============================================================================

def apply_warp(image, flow_map):
    """Apply flow-based warp correction to an image."""
    h, w = image.shape[:2]
    fh, fw = flow_map.shape[:2]
    
    # Scale flow to image resolution
    flow = cv2.resize(flow_map, (w, h), interpolation=cv2.INTER_CUBIC)
    flow[:,:,0] *= (w / fw)
    flow[:,:,1] *= (h / fh)
    
    # Build inverse remap coordinates
    ys, xs = np.mgrid[0:h, 0:w].astype(np.float32)
    map_x = xs - flow[:,:,0]
    map_y = ys - flow[:,:,1]
    
    corrected = cv2.remap(image, map_x, map_y,
                          interpolation=cv2.INTER_CUBIC,
                          borderMode=cv2.BORDER_REFLECT_101)
    return corrected


def run_baseline(test_images, flow_map):
    """Apply warp map to all test images."""
    print("\n" + "=" * 60)
    print("PHASE 2: BASELINE CORRECTION")
    print("=" * 60)
    
    start_time = time.time()
    processed = 0
    
    for img_path in test_images:
        try:
            image = cv2.imread(str(img_path))
            if image is None:
                continue
            
            corrected = apply_warp(image, flow_map)
            cv2.imwrite(str(BASELINE_DIR / img_path.name), corrected,
                       [cv2.IMWRITE_JPEG_QUALITY, JPEG_QUALITY])
            processed += 1
            
            if processed % 100 == 0:
                elapsed = time.time() - start_time
                print(f"  [{processed}/{len(test_images)}] {elapsed:.0f}s elapsed")
        except Exception as e:
            print(f"  Error on {img_path.name}: {e}")
    
    elapsed = time.time() - start_time
    print(f"  Baseline complete: {processed}/{len(test_images)} in {elapsed:.0f}s")
    return processed

# ============================================================================
# PHASE 3: TEST-TIME OPTIMIZATION
# ============================================================================

def apply_radial_correction(image, k1, k2, cx_off=0.0, cy_off=0.0):
    """Apply radial distortion correction using OpenCV undistort."""
    h, w = image.shape[:2]
    fx = fy = np.sqrt(w**2 + h**2)
    cx = w / 2.0 + cx_off * w
    cy = h / 2.0 + cy_off * h
    
    K = np.array([[fx, 0, cx], [0, fy, cy], [0, 0, 1]], dtype=np.float64)
    dist = np.array([k1, k2, 0, 0, 0], dtype=np.float64)
    
    new_K, roi = cv2.getOptimalNewCameraMatrix(K, dist, (w, h), 1, (w, h))
    return cv2.undistort(image, K, dist, None, new_K)


def detect_lines(image):
    """Detect line segments, return (lines_list, total_length)."""
    gray = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY) if len(image.shape) == 3 else image
    lsd = cv2.createLineSegmentDetector(cv2.LSD_REFINE_STD)
    lines, _, _, _ = lsd.detect(gray)
    
    if lines is None:
        return [], 0.0
    
    result = []
    total = 0.0
    for l in lines:
        x1, y1, x2, y2 = l[0]
        length = np.sqrt((x2-x1)**2 + (y2-y1)**2)
        if length >= MIN_LINE_LENGTH:
            angle = np.arctan2(y2-y1, x2-x1)
            result.append((length, angle))
            total += length
    return result, total


def line_score(image):
    """Score image by straight-line content. Higher = better."""
    lines, total_length = detect_lines(image)
    if not lines:
        return 0.0, 0
    
    h, w = image.shape[:2]
    diag = np.sqrt(h**2 + w**2)
    length_score = total_length / diag
    
    # Angular concentration near H/V
    angles = np.array([a for _, a in lines]) % np.pi
    lengths = np.array([l for l, _ in lines])
    
    dist_h = np.minimum(angles, np.pi - angles)
    dist_v = np.abs(angles - np.pi/2)
    dist_cardinal = np.minimum(dist_h, dist_v)
    
    wsum = np.sum(lengths)
    alignment = 1.0 - np.sum(dist_cardinal * lengths) / (wsum * np.pi/4) if wsum > 0 else 0
    
    return length_score * (0.5 + 0.5 * alignment), len(lines)


def optimize_objective(params, image):
    """Objective for optimizer (minimized)."""
    try:
        corrected = apply_radial_correction(image, *params)
        s, _ = line_score(corrected)
        return -s
    except:
        return 0.0


def optimize_single_image(args):
    """Optimize one image. Returns (name, params, score, n_lines, status)."""
    img_path, init_params = args
    name = Path(img_path).name
    
    try:
        image = cv2.imread(str(img_path))
        if image is None:
            return name, init_params, 0.0, 0, "load_fail"
        
        # Downsize for speed
        h, w = image.shape[:2]
        scale = OPTIM_RESOLUTION / max(h, w)
        if scale < 1.0:
            image = cv2.resize(image, (int(w*scale), int(h*scale)), 
                             interpolation=cv2.INTER_AREA)
        
        # Initial score
        init_score, init_lines = line_score(image)
        
        if init_lines < MIN_LINES_FOR_OPTIM:
            return name, init_params, init_score, init_lines, "few_lines"
        
        # Optimize
        result = minimize(
            optimize_objective,
            x0=init_params,
            args=(image,),
            method='Nelder-Mead',
            options={'maxiter': MAX_ITERATIONS, 'xatol': 1e-4, 'fatol': 1e-4, 'adaptive': True}
        )
        
        best_params = result.x.tolist()
        best_score = -result.fun
        
        corrected = apply_radial_correction(image, *best_params)
        _, best_lines = line_score(corrected)
        
        if best_score > init_score * 1.05:
            return name, best_params, best_score, best_lines, "optimized"
        else:
            return name, init_params, init_score, init_lines, "no_improvement"
    
    except Exception as e:
        return name, init_params, 0.0, 0, f"error:{e}"


def estimate_init_params(flow_map):
    """Estimate initial distortion params from warp map."""
    h, w = flow_map.shape[:2]
    cy, cx = h/2.0, w/2.0
    ys, xs = np.mgrid[0:h, 0:w].astype(np.float64)
    
    rx = (xs - cx) / w
    ry = (ys - cy) / h
    r = np.sqrt(rx**2 + ry**2)
    r_mag = r + 1e-8
    
    radial_disp = (flow_map[:,:,0] * rx + flow_map[:,:,1] * ry) / r_mag
    diag = np.sqrt(w**2 + h**2)
    radial_disp_norm = radial_disp / diag
    
    mask = r > 0.1
    r_flat = r[mask].flatten()
    d_flat = radial_disp_norm[mask].flatten()
    
    if len(r_flat) > 5000:
        idx = np.random.choice(len(r_flat), 5000, replace=False)
        r_flat, d_flat = r_flat[idx], d_flat[idx]
    
    A = np.column_stack([r_flat**2, r_flat**4])
    result = np.linalg.lstsq(A, d_flat, rcond=None)
    k1_est, k2_est = result[0]
    
    print(f"  Estimated initial params: k1={k1_est:.6f}, k2={k2_est:.6f}")
    return [float(k1_est), float(k2_est), 0.0, 0.0]


def run_optimizer(test_images, flow_map):
    """Run per-image optimization on all test images."""
    print("\n" + "=" * 60)
    print("PHASE 3: TEST-TIME OPTIMIZATION")
    print(f"  {NUM_WORKERS} workers, {MAX_ITERATIONS} iterations/image")
    print("=" * 60)
    
    start_time = time.time()
    
    # Get initial params from warp map
    if flow_map is not None:
        init_params = estimate_init_params(flow_map)
    else:
        init_params = [0.0, 0.0, 0.0, 0.0]
    
    # Prepare args
    args_list = [(str(img), init_params) for img in test_images]
    
    # Run optimization
    results = {}
    status_counts = {}
    
    if NUM_WORKERS > 1:
        with ProcessPoolExecutor(max_workers=NUM_WORKERS) as executor:
            futures = {executor.submit(optimize_single_image, a): a for a in args_list}
            done = 0
            for future in as_completed(futures):
                name, params, score, n_lines, status = future.result()
                results[name] = {'params': params, 'score': score, 
                               'num_lines': n_lines, 'status': status}
                status_counts[status] = status_counts.get(status, 0) + 1
                done += 1
                if done % 50 == 0:
                    elapsed = time.time() - start_time
                    print(f"  [{done}/{len(test_images)}] {elapsed:.0f}s | {status_counts}")
    else:
        for i, args in enumerate(args_list):
            name, params, score, n_lines, status = optimize_single_image(args)
            results[name] = {'params': params, 'score': score,
                           'num_lines': n_lines, 'status': status}
            status_counts[status] = status_counts.get(status, 0) + 1
            if (i+1) % 50 == 0:
                elapsed = time.time() - start_time
                print(f"  [{i+1}/{len(test_images)}] {elapsed:.0f}s | {status_counts}")
    
    print(f"\n  Optimization results: {status_counts}")
    print(f"  Time: {(time.time()-start_time)/60:.1f} minutes")
    
    # Apply corrections at full resolution
    print(f"\n  Applying optimized corrections at full resolution...")
    applied = 0
    
    for img_path in test_images:
        try:
            name = img_path.name
            image = cv2.imread(str(img_path))
            if image is None:
                continue
            
            if name in results and results[name]['params'] is not None:
                params = results[name]['params']
                if results[name]['status'] == 'optimized':
                    corrected = apply_radial_correction(image, *params)
                elif flow_map is not None:
                    corrected = apply_warp(image, flow_map)
                else:
                    corrected = image
            elif flow_map is not None:
                corrected = apply_warp(image, flow_map)
            else:
                corrected = image
            
            cv2.imwrite(str(OPTIMIZED_DIR / name), corrected,
                       [cv2.IMWRITE_JPEG_QUALITY, JPEG_QUALITY])
            applied += 1
        except:
            pass
    
    print(f"  Applied: {applied}/{len(test_images)}")
    
    # Save results JSON
    with open(OPTIMIZED_DIR / "results.json", 'w') as f:
        json.dump(results, f, indent=2)
    
    return results

# ============================================================================
# PHASE 4: ENSEMBLE BLENDING
# ============================================================================

def blend_weight(num_lines):
    """Compute optimizer weight based on line count."""
    if num_lines < MIN_LINES_ANY_OPTIM:
        return 0.0
    if num_lines >= MIN_LINES_FULL_OPTIM:
        return OPTIM_WEIGHT_MAX
    t = (num_lines - MIN_LINES_ANY_OPTIM) / (MIN_LINES_FULL_OPTIM - MIN_LINES_ANY_OPTIM)
    return OPTIM_WEIGHT_BASE + t * (OPTIM_WEIGHT_MAX - OPTIM_WEIGHT_BASE)


def run_ensemble(test_images, optim_results):
    """Blend baseline and optimized results."""
    print("\n" + "=" * 60)
    print("PHASE 4: ENSEMBLE BLENDING")
    print("=" * 60)
    
    start_time = time.time()
    stats = {'full_base': 0, 'full_opt': 0, 'blended': 0, 'base_only': 0}
    
    for img_path in test_images:
        name = img_path.name
        base_path = BASELINE_DIR / name
        opt_path = OPTIMIZED_DIR / name
        out_path = ENSEMBLE_DIR / name
        
        try:
            has_base = base_path.exists()
            has_opt = opt_path.exists()
            
            if not has_base and not has_opt:
                continue
            
            if has_base and not has_opt:
                img = cv2.imread(str(base_path))
                cv2.imwrite(str(out_path), img, [cv2.IMWRITE_JPEG_QUALITY, JPEG_QUALITY])
                stats['base_only'] += 1
                continue
            
            if not has_base and has_opt:
                img = cv2.imread(str(opt_path))
                cv2.imwrite(str(out_path), img, [cv2.IMWRITE_JPEG_QUALITY, JPEG_QUALITY])
                stats['full_opt'] += 1
                continue
            
            # Both exist â€” blend
            img_base = cv2.imread(str(base_path))
            img_opt = cv2.imread(str(opt_path))
            
            if img_base is None or img_opt is None:
                continue
            
            # Get line count
            n_lines = 0
            if name in optim_results:
                n_lines = optim_results[name].get('num_lines', 0)
            
            w_opt = blend_weight(n_lines)
            
            if w_opt == 0.0:
                output = img_base
                stats['full_base'] += 1
            elif w_opt >= OPTIM_WEIGHT_MAX:
                output = img_opt
                stats['full_opt'] += 1
            else:
                if img_base.shape != img_opt.shape:
                    img_opt = cv2.resize(img_opt, (img_base.shape[1], img_base.shape[0]))
                output = cv2.addWeighted(img_opt, w_opt, img_base, 1-w_opt, 0)
                stats['blended'] += 1
            
            cv2.imwrite(str(out_path), output, [cv2.IMWRITE_JPEG_QUALITY, JPEG_QUALITY])
        except:
            pass
    
    print(f"  Blend stats: {stats}")
    print(f"  Time: {(time.time()-start_time):.0f}s")

# ============================================================================
# ZIP & SUBMIT
# ============================================================================

def create_zip(output_dir, zip_name):
    """Create submission zip from output directory."""
    zip_path = WORKING_DIR / zip_name
    images = sorted(output_dir.glob("*.jpg"))
    
    print(f"\n  Zipping {len(images)} images...")
    with zipfile.ZipFile(str(zip_path), 'w', zipfile.ZIP_DEFLATED) as zf:
        for img in images:
            zf.write(str(img), img.name)
    
    size = zip_path.stat().st_size / (1024*1024)
    print(f"  Created: {zip_path} ({size:.1f} MB)")
    print(f"  Upload to: https://bounty.autohdr.com")
    return zip_path

# ============================================================================
# MAIN PIPELINE
# ============================================================================

def main():
    total_start = time.time()
    
    print("=" * 70)
    print("AUTOMATIC LENS CORRECTION â€” FULL PIPELINE")
    print("=" * 70)
    
    # Phase 0: Discover data
    train_pairs, test_images = discover_data()
    
    if not train_pairs:
        print("\nFATAL: No training pairs found!")
        print(f"Expected at: {TRAIN_DIR}")
        # List what's actually in the competition dir
        if COMPETITION_DIR.exists():
            print(f"\nContents of {COMPETITION_DIR}:")
            for item in sorted(COMPETITION_DIR.iterdir()):
                print(f"  {item.name}/ " if item.is_dir() else f"  {item.name}")
        return
    
    if not test_images:
        print("\nFATAL: No test images found!")
        print(f"Expected at: {TEST_DIR}")
        if COMPETITION_DIR.exists():
            print(f"\nContents of {COMPETITION_DIR}:")
            for item in sorted(COMPETITION_DIR.iterdir()):
                print(f"  {'[DIR] ' + item.name if item.is_dir() else item.name}")
                if item.is_dir():
                    sub = list(item.iterdir())[:5]
                    for s in sub:
                        print(f"    {s.name}")
                    if len(list(item.iterdir())) > 5:
                        print(f"    ... and {len(list(item.iterdir()))-5} more")
        return
    
    # Phase 1: Extract warp map
    flow_map = extract_warp_map(train_pairs)
    if flow_map is None:
        print("FATAL: Could not extract warp map!")
        return
    
    # Phase 2: Baseline
    run_baseline(test_images, flow_map)
    
    # Create baseline zip (safety net)
    create_zip(BASELINE_DIR, "submission_baseline.zip")
    
    # Phase 3: Optimize
    optim_results = run_optimizer(test_images, flow_map)
    
    # Create optimized zip  
    create_zip(OPTIMIZED_DIR, "submission_optimized.zip")
    
    # Phase 4: Ensemble
    run_ensemble(test_images, optim_results)
    
    # Create final zip
    create_zip(ENSEMBLE_DIR, "submission_ensemble.zip")
    
    # Summary
    total = time.time() - total_start
    print(f"\n{'=' * 70}")
    print(f"PIPELINE COMPLETE â€” Total time: {total/60:.1f} minutes")
    print(f"{'=' * 70}")
    print(f"\nSubmission files in /kaggle/working/:")
    print(f"  submission_baseline.zip   â€” Upload for Phase 2 score")
    print(f"  submission_optimized.zip  â€” Upload for Phase 3 score")  
    print(f"  submission_ensemble.zip   â€” Upload for BEST score (Phase 4)")
    print(f"\nUpload any zip to: https://bounty.autohdr.com")
    print(f"Then download submission.csv and submit to Kaggle")


if __name__ == "__main__":
    main()
