Source code for BsplineQuantRegpy.examples.comparison_example

#!/usr/bin/env python
# -*- coding: utf-8 -*-

"""
Script de comparaison des différents degrés de splines avec contraintes.
Peut être exécuté indépendamment ou appelé depuis la GUI.
"""

import numpy as np
import matplotlib.pyplot as plt
import sys
import os
import warnings
fie=__file__
# Ajouter le chemin du projet si exécuté indépendamment
print(__name__)
print(__file__)

if __name__ == "__main__":
    PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath("./")))
    SRC_PATH = os.path.join(PROJECT_ROOT, 'src')
    if SRC_PATH not in sys.path:
        sys.path.insert(0, SRC_PATH)


from BsplineQuantRegpy import (
        SplineLinearQuant,
        SplineQuadraticQuant,
        SplineCubicQuant,
        SplineQuarticQuant
    )

    


def generate_comparison_data(n=300, seed=42):
    """
    Génère des données de test pour la comparaison.
    
    Returns
    -------
    x : array
        Données x
    y : array
        Données y
    knots : array
        Nœuds pour les splines
    """
    np.random.seed(seed)
    x = np.linspace(0, 1, n)
    
    # Fonction avec différentes formes selon les régions
    y = np.zeros(n)
    
    # Région 1: [0, 0.3] - croissante convexe
    idx1 = x <= 0.3
    y[idx1] = 2 * x[idx1]**2 + 0.2 * np.random.randn(np.sum(idx1))
    
    # Région 2: [0.3, 0.6] - croissante concave
    idx2 = (x > 0.3) & (x <= 0.6)
    x2 = x[idx2]
    y[idx2] = 0.18 + 0.5 * np.sqrt(x2 - 0.3) + 0.2 * np.random.randn(np.sum(idx2))
    
    # Région 3: [0.6, 1] - décroissante
    idx3 = x > 0.6
    x3 = x[idx3]
    y[idx3] = 0.7 - 0.8 * (x3 - 0.6)**2 + 0.2 * np.random.randn(np.sum(idx3))
    
    # Nœuds
    kn = 15
    knots = np.quantile(x, np.linspace(0, 1, kn + 1))
    
    return x, y, knots


def run_comparison_analysis(x, y, knots, tau=0.5, solver='CLARABEL', 
                            show_plots=True, return_results=False):
    """
    Exécute un test de comparaison des degrés.
    
    Parameters
    ----------
    x, y : array
        Données
    knots : array
        Nœuds
    tau : float
        Quantile
    solver : str
        Solveur à utiliser
    show_plots : bool
        Afficher les graphiques
    return_results : bool
        Retourner les résultats
    
    Returns
    -------
    dict or None
        Résultats si return_results=True
    """
    x_eval = np.linspace(0, 1, 500)
    results = {}
    
    # ============ 1. Sans contraintes ============
    print("\n--- Sans contraintes ---")
    results['none'] = {}
    colors = ['blue', 'green', 'red', 'orange']
    deg_labels = ['Linéaire (deg 1)', 'Quadratique (deg 2)', 
                  'Cubique (deg 3)', 'Quartique (deg 4)']
    
    for d, color, label in zip([1, 2, 3, 4], colors, deg_labels):
        try:
            if d == 1:
                res = SplineLinearQuant(x, y, knots, tau=tau, 
                                           monot=0, solver=solver)
            elif d == 2:
                res = SplineQuadraticQuant(x, y, knots, tau=tau,
                                              monot=0, cv=0, solver=solver)
            elif d == 3:
                res = SplineCubicQuant(x, y, knots, tau=tau,
                                          monot=0, cv=0, der3=0, solver=solver)
            else:
                res = SplineQuarticQuant(x, y, knots, tau=tau,
                                            monot=0, cv=0, d3=0, solver=solver)
            
            if res is not None:
                results['none'][d] = {
                    'spline': res,
                    'color': color,
                    'label': label,
                    'y_eval': res(x_eval)
                }
                print(f"  ✓ {label}")
            else:
                print(f"  ✗ {label} (échec)")
        except Exception as e:
            print(f"  ✗ {label}: {e}")
    
    # ============ 2. Contrainte croissante ============
    print("\n--- Contrainte croissante ---")
    results['increasing'] = {}
    
    for d, color, label in zip([1, 2, 3, 4], colors, deg_labels):
        try:
            if d == 1:
                res = SplineLinearQuant(x, y, knots, tau=tau, 
                                           monot=1, solver=solver)
            elif d == 2:
                res = SplineQuadraticQuant(x, y, knots, tau=tau,
                                              monot=1, cv=0, solver=solver)
            elif d == 3:
                res = SplineCubicQuant(x, y, knots, tau=tau,
                                          monot=1, cv=0, der3=0, solver=solver)
            else:
                res = SplineQuarticQuant(x, y, knots, tau=tau,
                                            monot=1, cv=0, d3=0, solver=solver)
            
            if res is not None:
                results['increasing'][d] = {
                    'spline': res,
                    'color': color,
                    'label': label,
                    'y_eval': res(x_eval)
                }
                print(f"  ✓ {label}")
            else:
                print(f"  ✗ {label} (échec)")
        except Exception as e:
            print(f"  ✗ {label}: {e}")
    
    # ============ 3. Contrainte convexe ============
    print("\n--- Contrainte convexe ---")
    results['convex'] = {}
    deg2_labels = ['Quadratique (deg 2)', 'Cubique (deg 3)', 'Quartique (deg 4)']
    colors2 = ['green', 'red', 'orange']
    
    for d, color, label in zip([2, 3, 4], colors2, deg2_labels):
        try:
            if d == 2:
                res = SplineQuadraticQuant(x, y, knots, tau=tau,
                                              monot=0, cv=1, solver=solver)
            elif d == 3:
                res = SplineCubicQuant(x, y, knots, tau=tau,
                                          monot=0, cv=1, der3=0, solver=solver)
            else:
                res = SplineQuarticQuant(x, y, knots, tau=tau,
                                            monot=0, cv=1, d3=0, solver=solver)
            
            if res is not None:
                results['convex'][d] = {
                    'spline': res,
                    'color': color,
                    'label': label,
                    'y_eval': res(x_eval)
                }
                print(f"  ✓ {label}")
            else:
                print(f"  ✗ {label} (échec)")
        except Exception as e:
            print(f"  ✗ {label}: {e}")
    
    # ============ 4. Contrainte croissante + convexe ============
    print("\n--- Contrainte croissante + convexe ---")
    results['inc_convex'] = {}
    
    for d, color, label in zip([3, 4], ['red', 'orange'], 
                               ['Cubique (deg 3)', 'Quartique (deg 4)']):
        try:
            if d == 3:
                res = SplineCubicQuant(x, y, knots, tau=tau,
                                          monot=1, cv=1, der3=0, solver=solver)
            else:
                res = SplineQuarticQuant(x, y, knots, tau=tau,
                                            monot=1, cv=1, d3=0, solver=solver)
            
            if res is not None:
                results['inc_convex'][d] = {
                    'spline': res,
                    'color': color,
                    'label': label,
                    'y_eval': res(x_eval)
                }
                print(f"  ✓ {label}")
            else:
                print(f"  ✗ {label} (échec)")
        except Exception as e:
            print(f"  ✗ {label}: {e}")
    
    # Ajouter les références linéaire et quadratique
    try:
        res_lin = SplineLinearQuant(x, y, knots, tau=tau, monot=1, solver=solver)
        if res_lin is not None:
            results['inc_convex']['linear_ref'] = {
                'spline': res_lin,
                'color': 'blue',
                'label': 'Linéaire ',
                'y_eval': res_lin(x_eval),
                'is_ref': True
            }
    except:
        pass
    
    try:
        res_quad = SplineQuadraticQuant(x, y, knots, tau=tau, 
                                           monot=1, cv=0, solver=solver)
        if res_quad is not None:
            results['inc_convex']['quad_ref'] = {
                'spline': res_quad,
                'color': 'green',
                'label': 'Quadratique ',
                'y_eval': res_quad(x_eval),
                'is_ref': True
            }
    except:
        pass
    
    # ============ Affichage ============
    if show_plots:
        fig, axes = plt.subplots(2, 2, figsize=(14, 10))
        
        # 1. Sans contraintes
        ax = axes[0, 0]
        ax.scatter(x, y, alpha=0.3, s=10, color='gray', label='Données')
        for d in results['none']:
            data = results['none'][d]
            ax.plot(x_eval, data['y_eval'], color=data['color'], 
                   linewidth=2, label=data['label'])
        ax.plot(knots, np.ones_like(knots)*max(y)*0.95, 'k|', markersize=8, label='Nœuds')
        ax.set_xlabel('x')
        ax.set_ylabel('y')
        ax.set_title('Sans contraintes')
        ax.legend(fontsize='small')
        ax.grid(True, alpha=0.3)
        
        # 2. Contrainte croissante
        ax = axes[0, 1]
        ax.scatter(x, y, alpha=0.3, s=10, color='gray', label='Données')
        for d in results['increasing']:
            data = results['increasing'][d]
            ax.plot(x_eval, data['y_eval'], color=data['color'], 
                   linewidth=2, label=data['label'])
        ax.plot(knots, np.ones_like(knots)*max(y)*0.95, 'k|', markersize=8, label='Nœuds')
        ax.set_xlabel('x')
        ax.set_ylabel('y')
        ax.set_title('Contrainte croissante')
        ax.legend(fontsize='small')
        ax.grid(True, alpha=0.3)
        
        # 3. Contrainte convexe
        ax = axes[1, 0]
        ax.scatter(x, y, alpha=0.3, s=10, color='gray', label='Données')
        for d in results['convex']:
            data = results['convex'][d]
            ax.plot(x_eval, data['y_eval'], color=data['color'], 
                   linewidth=2, label=data['label'])
        ax.plot(knots, np.ones_like(knots)*max(y)*0.95, 'k|', markersize=8, label='Nœuds')
        ax.set_xlabel('x')
        ax.set_ylabel('y')
        ax.set_title('Contrainte convexe')
        ax.legend(fontsize='small')
        ax.grid(True, alpha=0.3)
        
        # 4. Contrainte croissante + convexe
        ax = axes[1, 1]
        ax.scatter(x, y, alpha=0.3, s=10, color='gray', label='Données')
        for key in results['inc_convex']:
            data = results['inc_convex'][key]
            linestyle = '--' if data.get('is_ref', False) else '-'
            linewidth = 1.5 if data.get('is_ref', False) else 2
            ax.plot(x_eval, data['y_eval'], color=data['color'], 
                   linewidth=linewidth, linestyle=linestyle, label=data['label'])
        ax.plot(knots, np.ones_like(knots)*max(y)*0.95, 'k|', markersize=8, label='Nœuds')
        ax.set_xlabel('x')
        ax.set_ylabel('y')
        ax.set_title('Contrainte croissante + convexe')
        ax.legend(fontsize='small')
        ax.grid(True, alpha=0.3)
        
        plt.tight_layout()
        plt.show()
    
    if return_results:
        return results
    else:
        return None


def main():
    """Fonction principale pour exécution indépendante."""
    print("=" * 70)
    print("COMPARAISON DES DEGRÉS DE SPLINES AVEC CONTRAINTES")
    print("=" * 70)
    
    # Générer les données
    x, y, knots = generate_comparison_data()
    
    # Exécuter l'analyse
    run_comparison_analysis(x, y, knots, tau=0.5, solver='CLARABEL', show_plots=True)
    
    print("\n" + "=" * 70)
    print("FIN DE LA COMPARAISON")
    print("=" * 70)


if __name__ == "__main__":
    main()