"""
RECONSTRUCTION 3D CHAMBRE
Pipeline complet : filtre IRL → depth MiDaS → SfM → export .ply

DÉPENDANCES :
    pip install opencv-python numpy scikit-learn joblib torch torchvision timm matplotlib pyntcloud pandas scipy

USAGE :
    python reconstruct_3d.py

NOTES :
    - Lance train_classifier.py d'abord pour générer le classifieur
    - Le modèle MiDaS (~400MB) se télécharge automatiquement au premier lancement
    - Résultat final : chambre_3d.ply → ouvre dans MeshLab
"""

import cv2
import numpy as np
import os
import sys
import glob
import re
import joblib
import torch
import pandas as pd
from pyntcloud import PyntCloud
from scipy.spatial import cKDTree
from datetime import datetime, timedelta, timezone
from multiprocessing import Pool, cpu_count, freeze_support
import multiprocessing as mp
import json

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

# Classifieur IRL vs gameplay
CLASSIFIER_FILE = "classifier_irl_vs_gameplay.joblib"
SCALER_FILE     = "scaler_irl_vs_gameplay.joblib"

# Seuil de confiance IRL (0.0 à 1.0)
IRL_CONFIDENCE = 0.75

# Pas d'échantillonnage en secondes
SAMPLE_STEP = 5

# Nombre max de frames à traiter (None = toutes)
MAX_FRAMES = 2000

# Nb de workers parallèles — limité pour économiser la RAM
# Chaque worker charge une copie du classifieur + frame en mémoire
NB_WORKERS = min(4, max(1, cpu_count() - 2))

# Taille de traitement MiDaS
MIDAS_SIZE = (384, 384)

# Fichiers de sortie
OUTPUT_PLY       = "chambre_3d.ply"
OUTPUT_FRAMES    = "frames_irl/"
DEPTH_CACHE      = "depth_cache.npz"
FEATURES_CACHE   = "sfm_features.json"

# ============================================================
#  CHARGEMENT DU CLASSIFIEUR
# ============================================================

def load_classifier():
    if not os.path.exists(CLASSIFIER_FILE):
        print(f"[ERR] Classifieur introuvable : {CLASSIFIER_FILE}")
        print("      Lance d'abord : python train_classifier.py")
        sys.exit(1)
    clf    = joblib.load(CLASSIFIER_FILE)
    scaler = joblib.load(SCALER_FILE)
    print(f"[OK] Classifieur chargé")
    return clf, scaler


# ============================================================
#  CHARGEMENT MIDAS
# ============================================================

def load_midas():
    print("\n[MIDAS] Chargement du modèle de profondeur...")
    print("        (téléchargement ~400MB au premier lancement)")

    model_type = "DPT_Large"
    midas = torch.hub.load("intel-isl/MiDaS", model_type)
    midas.eval()

    # CPU uniquement
    device = torch.device("cpu")
    midas.to(device)

    transforms = torch.hub.load("intel-isl/MiDaS", "transforms")
    transform  = transforms.dpt_transform

    print(f"[OK] MiDaS chargé sur CPU")
    return midas, transform, device


# ============================================================
#  EXTRACTION FEATURES CLASSIFIEUR (même que train_classifier.py)
# ============================================================

def extract_features(frame: np.ndarray) -> np.ndarray:
    features = []
    frame_small = cv2.resize(frame, (128, 72))

    hsv = cv2.cvtColor(frame_small, cv2.COLOR_BGR2HSV)
    for i in range(3):
        hist = cv2.calcHist([hsv], [i], None, [16], [0, 256])
        hist = cv2.normalize(hist, hist).flatten()
        features.extend(hist.tolist())

    gray = cv2.cvtColor(frame_small, cv2.COLOR_BGR2GRAY).astype(float)
    lap  = cv2.Laplacian(gray, cv2.CV_64F)
    features.append(float(lap.var()))
    features.append(float(np.mean(np.abs(lap))))

    fft       = np.fft.fft2(gray)
    fft_shift = np.fft.fftshift(fft)
    magnitude = np.abs(fft_shift)
    h, w      = magnitude.shape
    hf_ring   = magnitude.copy()
    hf_ring[h//4:3*h//4, w//4:3*w//4] = 0
    features.append(float(np.mean(hf_ring)))
    features.append(float(np.std(hf_ring)))

    block_vars = []
    bsize = 8
    for y in range(0, gray.shape[0] - bsize, bsize):
        for x in range(0, gray.shape[1] - bsize, bsize):
            block = gray[y:y+bsize, x:x+bsize]
            block_vars.append(float(block.var()))
    features.append(float(np.mean(block_vars)))
    features.append(float(np.std(block_vars)))
    features.append(float(np.percentile(block_vars, 90)))

    sobelx  = cv2.Sobel(gray, cv2.CV_64F, 1, 0, ksize=3)
    sobely  = cv2.Sobel(gray, cv2.CV_64F, 0, 1, ksize=3)
    grad    = np.sqrt(sobelx**2 + sobely**2)
    features.append(float(np.mean(grad)))
    features.append(float(np.std(grad)))
    features.append(float(np.percentile(grad, 95)))

    sat = hsv[:, :, 1].astype(float)
    features.append(float(np.mean(sat)))
    features.append(float(np.std(sat)))

    edges_thin  = cv2.Canny(frame_small, 100, 200)
    edges_thick = cv2.Canny(frame_small, 30, 80)
    thin_count  = float(np.sum(edges_thin > 0))
    thick_count = float(np.sum(edges_thick > 0))
    ratio       = thin_count / (thick_count + 1e-6)
    features.append(ratio)
    features.append(thin_count)
    features.append(thick_count)

    return np.array(features, dtype=float)


# ============================================================
#  ÉTAPE 1 — EXTRACTION DES FRAMES IRL
# ============================================================

def is_irl_frame(frame, clf, scaler) -> tuple:
    """Retourne (is_irl: bool, confidence: float)"""
    feats  = extract_features(frame).reshape(1, -1)
    scaled = scaler.transform(feats)
    # Force execution sequentielle pour eviter Loky dans un worker multiprocessing
    import joblib
    with joblib.parallel_backend("threading", n_jobs=1):
        proba = clf.predict_proba(scaled)[0]
    # classe 1 = IRL
    classes = list(clf.classes_)
    irl_idx = classes.index(1) if 1 in classes else 1
    conf    = float(proba[irl_idx])
    return conf >= IRL_CONFIDENCE, conf


def parse_date_from_filename(filepath: str) -> datetime:
    basename = os.path.basename(filepath)
    pattern  = r'(\d{4}-\d{2}-\d{2})_(\d{2}-\d{2}-\d{2})'
    match    = re.search(pattern, basename)
    if not match:
        return None
    date_str = match.group(1)
    time_str = match.group(2).replace('-', ':')
    dt = datetime.strptime(f"{date_str} {time_str}", "%Y-%m-%d %H:%M:%S")
    return dt.replace(tzinfo=timezone.utc)


# Variables globales dans chaque worker (initialisees via pool initializer)
_worker_clf    = None
_worker_scaler = None

def _worker_init(classifier_file, scaler_file):
    """Initialise le classifieur une seule fois par worker.
    Desactive tout parallelisme interne (sklearn/Loky/OpenMP/BLAS)
    pour eviter le warning 'Loky-backed parallel loops cannot be called
    in a multiprocessing context'.
    """
    import os
    # Desactive le parallelisme sklearn/joblib/Loky dans les workers
    os.environ["LOKY_MAX_CPU_COUNT"]    = "1"
    os.environ["OMP_NUM_THREADS"]       = "1"
    os.environ["MKL_NUM_THREADS"]       = "1"
    os.environ["OPENBLAS_NUM_THREADS"]  = "1"
    os.environ["NUMEXPR_NUM_THREADS"]   = "1"

    import joblib as _joblib
    global _worker_clf, _worker_scaler
    _worker_clf    = _joblib.load(classifier_file)
    _worker_scaler = _joblib.load(scaler_file)

    # Force sklearn a utiliser 1 seul thread (n_jobs=1) pour ce process
    try:
        from sklearn.utils import parallel_backend
        import threading
        # Patch global : toutes les operations sklearn dans ce worker seront sequentielles
        _worker_clf.n_jobs = 1
    except Exception:
        pass


def extract_irl_frames_from_video(args) -> list:
    """
    Worker multiprocessing : extrait les frames IRL d'une vidéo.
    Retourne liste de chemins de frames sauvegardées.
    """
    video_file, out_dir, max_per_video = args
    clf    = _worker_clf
    scaler = _worker_scaler

    if not os.path.exists(video_file):
        return []

    cap = cv2.VideoCapture(video_file)
    if not cap.isOpened():
        return []

    fps          = cap.get(cv2.CAP_PROP_FPS) or 30
    total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
    duration_sec = int(total_frames / fps)

    saved_frames = []
    curr         = 0
    prev_gray    = None
    video_base   = os.path.splitext(os.path.basename(video_file))[0]

    while curr <= duration_sec and len(saved_frames) < max_per_video:
        cap.set(cv2.CAP_PROP_POS_MSEC, curr * 1000)
        ret, frame = cap.read()
        if not ret:
            break

        # Filtre classifieur
        is_irl, conf = is_irl_frame(frame, clf, scaler)

        if is_irl:
            # Filtre variation brutale (passage devant cam)
            gray = cv2.cvtColor(frame, cv2.COLOR_BGR2GRAY)
            if prev_gray is not None:
                variation = float(np.mean(cv2.absdiff(prev_gray, gray)))
                if variation > 35.0:
                    prev_gray = gray.copy()
                    curr += SAMPLE_STEP
                    continue
            prev_gray = gray.copy() if prev_gray is None else prev_gray

            # Sauvegarde frame
            frame_name = f"{video_base}_{curr:06d}.jpg"
            frame_path = os.path.join(out_dir, frame_name)
            cv2.imwrite(frame_path, frame, [cv2.IMWRITE_JPEG_QUALITY, 85])
            saved_frames.append({
                "path": frame_path,
                "video": video_file,
                "t_sec": curr,
                "conf": conf,
            })

        curr += SAMPLE_STEP

    cap.release()
    return saved_frames


def extract_all_irl_frames(clf, scaler) -> list:
    """Lance l'extraction en parallèle sur tous les .mp4"""
    os.makedirs(OUTPUT_FRAMES, exist_ok=True)

    mp4_files = sorted(glob.glob("*.mp4"))
    if not mp4_files:
        print("[ERR] Aucun .mp4 trouvé")
        sys.exit(1)

    print(f"\n[EXTRACT] {len(mp4_files)} vidéos trouvées")
    print(f"[EXTRACT] {NB_WORKERS} workers parallèles")

    max_per_video = max(10, (MAX_FRAMES or 99999) // len(mp4_files))

    args_list = [
        (f, OUTPUT_FRAMES, max_per_video)
        for f in mp4_files
    ]

    # Multiprocessing sur les workers — context "spawn" pour eviter
    # les conflits avec PyTorch et joblib/Loky.
    # clf/scaler sont charges dans chaque worker via initializer (pas serialises dans args)
    ctx = mp.get_context("spawn")
    with ctx.Pool(
        processes=NB_WORKERS,
        initializer=_worker_init,
        initargs=(CLASSIFIER_FILE, SCALER_FILE),
    ) as pool:
        results = pool.map(extract_irl_frames_from_video, args_list)

    all_frames = []
    for r in results:
        all_frames.extend(r)

    print(f"\n[OK] {len(all_frames)} frames IRL extraites")

    if MAX_FRAMES and len(all_frames) > MAX_FRAMES:
        # Garde les frames les plus confiantes
        all_frames = sorted(all_frames, key=lambda x: -x["conf"])[:MAX_FRAMES]
        print(f"[OK] Réduit à {MAX_FRAMES} frames (plus confiantes)")

    return all_frames


# ============================================================
#  ÉTAPE 2 — ESTIMATION DE PROFONDEUR (MiDaS)
# ============================================================

def estimate_depth(frame_bgr: np.ndarray, midas, transform, device) -> np.ndarray:
    """Retourne une carte de profondeur normalisée [0,1]."""
    img_rgb = cv2.cvtColor(frame_bgr, cv2.COLOR_BGR2RGB)
    input_batch = transform(img_rgb).to(device)

    with torch.no_grad():
        prediction = midas(input_batch)
        prediction = torch.nn.functional.interpolate(
            prediction.unsqueeze(1),
            size=frame_bgr.shape[:2],
            mode="bicubic",
            align_corners=False,
        ).squeeze()

    depth = prediction.cpu().numpy()

    # Normalisation [0, 1]
    dmin, dmax = depth.min(), depth.max()
    if dmax - dmin > 1e-6:
        depth = (depth - dmin) / (dmax - dmin)

    return depth.astype(np.float32)


DEPTH_CACHE_DIR = "depth_cache_frames/"

def _depth_cache_path(frame_path: str) -> str:
    """Retourne le chemin .npy du cache pour une frame."""
    basename = os.path.splitext(os.path.basename(frame_path))[0]
    return os.path.join(DEPTH_CACHE_DIR, basename + ".npy")


def compute_all_depths(frames: list, midas, transform, device) -> dict:
    """
    Calcule la carte de profondeur pour chaque frame.
    *** VERSION ÉCONOME EN MÉMOIRE ***
    Chaque depth map est sauvegardée individuellement sur disque.
    Retourne un dict {frame_path: chemin_npy} — chargé à la demande.
    """
    os.makedirs(DEPTH_CACHE_DIR, exist_ok=True)

    # Vérifie si tout est déjà en cache
    all_cached = all(os.path.exists(_depth_cache_path(f["path"])) for f in frames)
    if all_cached:
        print(f"\n[DEPTH] Cache complet trouvé dans {DEPTH_CACHE_DIR}")
        return {f["path"]: _depth_cache_path(f["path"]) for f in frames}

    print(f"\n[DEPTH] Estimation profondeur sur {len(frames)} frames (CPU)...")
    print("        (peut prendre du temps sur CPU — ~1-3s par frame)")

    depths_map = {}
    for i, frame_info in enumerate(frames):
        path       = frame_info["path"]
        cache_file = _depth_cache_path(path)

        if os.path.exists(cache_file):
            depths_map[path] = cache_file
            continue

        frame = cv2.imread(path)
        if frame is None:
            continue

        depth = estimate_depth(frame, midas, transform, device)
        np.save(cache_file, depth)    # sauvegarde immédiate
        depths_map[path] = cache_file # on garde le chemin, pas l'array
        del depth, frame              # libère la RAM

        if (i + 1) % 10 == 0:
            pct = 100 * (i + 1) / len(frames)
            print(f"  {i+1}/{len(frames)} ({pct:.0f}%)...", end="\r")

    print(f"\n[OK] Profondeurs calculées (stockées dans {DEPTH_CACHE_DIR})")
    return depths_map


# ============================================================
#  ÉTAPE 3 — STRUCTURE FROM MOTION (SfM léger)
# ============================================================

def compute_camera_poses(frames: list) -> list:
    """
    Estime les poses caméra entre frames consécutives via ORB + homographie.
    Retourne une liste de matrices de transformation 4x4.
    """
    print(f"\n[SFM] Calcul des poses caméra...")

    orb      = cv2.ORB_create(nfeatures=1000)
    matcher  = cv2.BFMatcher(cv2.NORM_HAMMING, crossCheck=True)

    poses    = [np.eye(4)]  # première pose = identité
    curr_pose = np.eye(4)

    # Paramètres caméra approximatifs pour 480p
    # Focale approximée pour un téléphone standard
    h, w  = 480, 854
    fx    = w * 0.8   # approximation focale
    fy    = fx
    cx, cy = w / 2, h / 2
    K     = np.array([[fx, 0, cx],
                      [0, fy, cy],
                      [0,  0,  1]], dtype=float)

    prev_frame = None
    prev_kp    = None
    prev_des   = None

    for i, frame_info in enumerate(frames[:-1]):
        frame = cv2.imread(frame_info["path"])
        if frame is None:
            poses.append(curr_pose.copy())
            continue

        gray = cv2.cvtColor(frame, cv2.COLOR_BGR2GRAY)
        kp, des = orb.detectAndCompute(gray, None)

        if prev_frame is not None and prev_des is not None and des is not None:
            matches = matcher.match(prev_des, des)
            matches = sorted(matches, key=lambda x: x.distance)

            if len(matches) >= 8:
                pts1 = np.float32([prev_kp[m.queryIdx].pt for m in matches])
                pts2 = np.float32([kp[m.trainIdx].pt for m in matches])

                # Matrice essentielle
                E, mask = cv2.findEssentialMat(
                    pts1, pts2, K,
                    method=cv2.RANSAC,
                    prob=0.999,
                    threshold=1.0
                )

                if E is not None:
                    _, R, t, _ = cv2.recoverPose(E, pts1, pts2, K, mask=mask)

                    # Construire transformation relative
                    T_rel      = np.eye(4)
                    T_rel[:3, :3] = R
                    T_rel[:3, 3]  = t.flatten()

                    curr_pose = curr_pose @ T_rel

        poses.append(curr_pose.copy())
        prev_frame = gray
        prev_kp    = kp
        prev_des   = des

        if (i + 1) % 50 == 0:
            print(f"  {i+1}/{len(frames)} poses...", end="\r")

    print(f"\n[OK] {len(poses)} poses calculées")
    return poses, K


# ============================================================
#  ÉTAPE 4 — FUSION NUAGE DE POINTS
# ============================================================

def build_point_cloud(frames: list, depths: dict, poses: list, K: np.ndarray) -> dict:
    """
    Fusionne les cartes de profondeur avec les poses caméra.
    *** VERSION ÉCONOME EN MÉMOIRE ***
    - Charge les depth maps depuis le disque une à une (lazy loading)
    - Accumule par chunks de 100 frames puis libère
    Retourne un dict {"points": np.ndarray (N,3), "colors": np.ndarray (N,3)}
    """
    print(f"\n[CLOUD] Construction du nuage de points...")

    fx, fy = K[0, 0], K[1, 1]
    cx, cy = K[0, 2], K[1, 2]

    # Grille de pixels (sous-échantillonné pour performance)
    skip = 4  # un pixel sur 4

    # Accumulation par chunks pour éviter le double-picage mémoire du vstack final
    CHUNK_SIZE = 100
    chunk_points = []
    chunk_colors = []
    merged_points = []
    merged_colors = []
    total_pts = 0

    def _flush_chunk():
        """Fusionne le chunk courant et libère."""
        nonlocal chunk_points, chunk_colors
        if chunk_points:
            merged_points.append(np.vstack(chunk_points))
            merged_colors.append(np.vstack(chunk_colors))
            chunk_points = []
            chunk_colors = []

    for i, frame_info in enumerate(frames):
        path = frame_info["path"]

        if path not in depths:
            continue

        # Lazy load depth depuis disque
        depth_src = depths[path]
        if isinstance(depth_src, str):
            depth = np.load(depth_src)
        else:
            depth = depth_src  # rétrocompatibilité si array direct

        frame = cv2.imread(path)
        if frame is None:
            del depth
            continue

        pose  = poses[i] if i < len(poses) else np.eye(4)

        h_f, w_f = depth.shape
        frame_rgb = cv2.cvtColor(
            cv2.resize(frame, (w_f, h_f)),
            cv2.COLOR_BGR2RGB
        ).astype(np.float32) / 255.0

        del frame  # libère la frame BGR

        # Pixels valides (profondeur non nulle)
        ys, xs = np.mgrid[0:h_f:skip, 0:w_f:skip]
        ys     = ys.flatten()
        xs     = xs.flatten()
        ds     = depth[ys, xs]
        del depth  # libère la depth map

        # Filtre profondeur nulle
        valid  = ds > 0.01
        ys, xs, ds = ys[valid], xs[valid], ds[valid]

        depth_scale = 3.0
        Z = ds * depth_scale
        X = (xs - cx) / fx * Z
        Y = (ys - cy) / fy * Z

        pts_cam   = np.stack([X, Y, Z, np.ones_like(Z)], axis=1)
        pts_world = (pose @ pts_cam.T).T[:, :3].astype(np.float32)
        colors    = frame_rgb[ys, xs].astype(np.float32)

        chunk_points.append(pts_world)
        chunk_colors.append(colors)
        total_pts += len(pts_world)

        # Flush tous les CHUNK_SIZE frames
        if (i + 1) % CHUNK_SIZE == 0:
            _flush_chunk()

        if (i + 1) % 20 == 0:
            print(f"  {i+1}/{len(frames)} frames | {total_pts:,} points...", end="\r")

    _flush_chunk()  # flush le dernier chunk

    if not merged_points:
        print("[ERR] Aucun point généré")
        sys.exit(1)

    all_points = np.vstack(merged_points)
    all_colors = np.vstack(merged_colors)
    del merged_points, merged_colors

    print(f"\n[OK] {len(all_points):,} points générés")
    return {"points": all_points, "colors": all_colors}


# ============================================================
#  ÉTAPE 5 — NETTOYAGE ET EXPORT
# ============================================================

def clean_and_export(pcd: dict):
    """
    Nettoie le nuage (supprime outliers, voxel downsampling) et exporte en .ply
    Utilise pyntcloud + numpy/scipy (pas besoin d'AVX/Open3D).
    """
    print(f"\n[CLEAN] Nettoyage du nuage de points...")

    points = pcd["points"]
    colors = pcd["colors"]

    print(f"  Avant nettoyage : {len(points):,} pts")

    # --- Suppression des outliers statistiques par batch ---
    k = 10  # réduit de 20 à 10 pour économiser la RAM
    tree = cKDTree(points)
    BATCH_O = 10000
    mean_dists = np.empty(len(points), dtype=np.float32)
    for start in range(0, len(points), BATCH_O):
        end = min(start + BATCH_O, len(points))
        dists_b, _ = tree.query(points[start:end], k=k + 1)
        mean_dists[start:end] = dists_b[:, 1:].mean(axis=1)
    threshold = mean_dists.mean() + 2.0 * mean_dists.std()
    mask = mean_dists < threshold
    del mean_dists
    points = points[mask]
    colors = colors[mask]
    print(f"  Après outliers  : {len(points):,} pts")

    # --- Voxel downsampling vectorisé (numpy unique) ---
    voxel_size = 0.02
    voxel_indices = np.floor(points / voxel_size).astype(np.int32)
    # Encode chaque voxel comme un entier unique pour np.unique
    vi_min = voxel_indices.min(axis=0)
    vi_shifted = voxel_indices - vi_min
    vi_max = vi_shifted.max(axis=0) + 1
    flat_idx = (vi_shifted[:, 0] * vi_max[1] * vi_max[2]
                + vi_shifted[:, 1] * vi_max[2]
                + vi_shifted[:, 2])
    _, first_occ = np.unique(flat_idx, return_index=True)
    points = points[first_occ]
    colors = colors[first_occ]
    print(f"  Après voxel     : {len(points):,} pts")

    # --- Estimation des normales (PCA locale vectorisée par batch) ---
    print("  Estimation des normales (batch vectorisé)...")
    normals = np.zeros_like(points)
    tree2   = cKDTree(points)
    k_nn    = 15  # k voisins fixes (plus rapide que query_ball_point)
    BATCH   = 2000  # traitement par batch pour limiter la RAM

    for start in range(0, len(points), BATCH):
        end = min(start + BATCH, len(points))
        batch = points[start:end]
        # k+1 car le point lui-même est inclus
        _, idxs_batch = tree2.query(batch, k=k_nn + 1)
        for j, idxs in enumerate(idxs_batch):
            neighbors = points[idxs[1:]]  # exclure le point lui-même
            if len(neighbors) < 3:
                continue
            cov = np.cov((neighbors - neighbors.mean(axis=0)).T)
            _, eigvecs = np.linalg.eigh(cov)
            normals[start + j] = eigvecs[:, 0]
        if (start // BATCH) % 5 == 0:
            pct = 100 * end / len(points)
            print(f"    normales {pct:.0f}%...", end="\r")
    print()

    # --- Export .ply via pyntcloud ---
    colors_uint8 = np.clip(colors * 255, 0, 255).astype(np.uint8)
    normals_f32  = normals.astype(np.float32)

    df = pd.DataFrame({
        "x": points[:, 0].astype(np.float32),
        "y": points[:, 1].astype(np.float32),
        "z": points[:, 2].astype(np.float32),
        "red":   colors_uint8[:, 0],
        "green": colors_uint8[:, 1],
        "blue":  colors_uint8[:, 2],
        "nx": normals_f32[:, 0],
        "ny": normals_f32[:, 1],
        "nz": normals_f32[:, 2],
    })

    cloud = PyntCloud(df)
    cloud.to_file(OUTPUT_PLY)

    print(f"\n[OK] Nuage exporté → {OUTPUT_PLY}")
    print(f"     Ouvre dans MeshLab : File → Import Mesh → {OUTPUT_PLY}")

    # Stats finales
    bbox = points.max(axis=0) - points.min(axis=0)
    print(f"\n  Dimensions reconstituées :")
    print(f"    Largeur   : {bbox[0]:.2f}m")
    print(f"    Profondeur: {bbox[1]:.2f}m")
    print(f"    Hauteur   : {bbox[2]:.2f}m")

    return {"points": points, "colors": colors}


# ============================================================
#  MAIN
# ============================================================

if __name__ == "__main__":
    # Nécessaire sur Windows et conseillé sur Linux avec PyTorch
    freeze_support()
    # "fork" est le défaut sur Linux mais cause des conflits avec PyTorch
    # On force "spawn" pour isoler proprement les workers du process PyTorch
    mp.set_start_method("spawn", force=True)

    print("=" * 60)
    print("  RECONSTRUCTION 3D CHAMBRE")
    print(f"  Workers : {NB_WORKERS} / {cpu_count()} cores disponibles")
    print("=" * 60)

    # 1. Classifieur
    clf, scaler = load_classifier()

    # 2. Extraction frames IRL (multiprocessing)
    print("\n[ÉTAPE 1/5] Extraction des frames IRL...")
    all_frames = extract_all_irl_frames(clf, scaler)

    if not all_frames:
        print("[ERR] Aucune frame IRL trouvée — vérifie le classifieur")
        sys.exit(1)

    # 3. MiDaS
    print("\n[ÉTAPE 2/5] Chargement MiDaS...")
    midas, transform, device = load_midas()

    # 4. Profondeur
    print("\n[ÉTAPE 3/5] Estimation de profondeur...")
    depths = compute_all_depths(all_frames, midas, transform, device)

    # 5. SfM
    print("\n[ÉTAPE 4/5] Structure from Motion...")
    poses, K = compute_camera_poses(all_frames)

    # 6. Nuage de points
    print("\n[ÉTAPE 5/5] Construction et export nuage 3D...")
    pcd = build_point_cloud(all_frames, depths, poses, K)
    pcd_final = clean_and_export(pcd)

    print("\n" + "=" * 60)
    print("  TERMINÉ !")
    print(f"  → Ouvre {OUTPUT_PLY} dans MeshLab")
    print("=" * 60)
