#!/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()