## Pre-script

# Script de traitement d'un signal bruité par moyenne glissante, filtre passe
# bas du premier et du second ordre

import numpy as np
from matplotlib import pyplot as plt

# Pour fermer toutes les anciennes figures au lancement du script
plt.close('all')

## Définition des fonctions pour tracer les figures de comparaison
# Tracer d'un plot simple
def plot_comparison(signal_brut, signal_filtre, methode_filtrage, params_filtrage, temps):
    """
    Trace le signal brut et le signal filtré, avec les étiquettes appropriées

    Paramètres :
    - signal_brut : ndarray (array numpy), signal brut
    - signal_filtre : ndarray (array numpy), signal filtré
    - methode_filtrage : str, nom de la méthode de filtrage (par exemple, 'Moyenne glissante', 'Filtre passe-bas premier ordre', etc.)
    - params_filtrage : str, paramètres spécifiques à la méthode de filtrage (par exemple, la taille de la fenêtre pour la moyenne glissante, ou la constante de temps pour le filtre passe-bas du premier ordre)
    - temps : ndarray (array numpy), vecteur de temps

    Retourne :
    - None
    """
    plt.figure(figsize=(10, 7))
    plt.plot(temps, signal_brut, label='Signal brut')
    plt.plot(temps, signal_filtre, label=f'Signal filtré ({methode_filtrage} - {params_filtrage})', linestyle='--')
    plt.xlabel('Temps (s)')
    plt.ylabel('Amplitude')
    plt.title(f'Comparaison du signal brut et du signal filtré ({methode_filtrage} - {params_filtrage})')
    plt.legend()
    plt.grid(True)
    plt.show()

# Tracer plusieurs plot sur la même figure
def plot_multiple_filter_comparison(signal_brut, signals_filtres, noms_filtrage, params_filtrage, temps):
    """
    Trace sur la même figure le signal brut et plusieurs signaux filtrés avec différentes paramètres,
    avec les étiquettes appropriées pour chaque méthode de filtrage et ses paramètres.

    Paramètres :
    - signal_brut : ndarray (array numpy), signal brut
    - signals_filtres : liste d'array numpy, liste des signaux filtrés à tracer
    - noms_filtrage : liste de str, noms des méthodes de filtrage correspondant à chaque signal filtré
    - params_filtrage : liste de str, paramètres spécifiques à chaque méthode de filtrage
    - temps : ndarray (array numpy), vecteur de temps

    Retourne :
    - None
    """
    plt.figure(figsize=(10, 7))
    plt.plot(temps, signal_brut, label='Signal brut')

    for i in range(len(signals_filtres)):
        plt.plot(temps, signals_filtres[i], label=f'Signal filtré ({noms_filtrage[i]} - {params_filtrage[i]})', linestyle='--')

    plt.xlabel('Temps (s)')
    plt.ylabel('Amplitude')
    plt.title('Comparaison du signal brut et des signaux filtrés')
    plt.legend()
    plt.grid(True)
    plt.show()

# Tracer de trois subplot
def plot_comparison_subplots(signal_brut, signal_filtre_base, nom_filtrage, param_base, temps):
    """
    Trace le signal brut et les signaux filtrés dans des subplots, avec les étiquettes appropriées.
    Utilise la même méthode de filtrage pour tous les subplots mais varie les paramètres.

    Paramètres :
    - signal_brut : ndarray (array numpy), signal brut
    - signal_filtre_base : ndarray (array numpy), signal filtré de référence pour une méthode de filtrage donnée
    - nom_filtrage : str, nom de la méthode de filtrage (par exemple, 'Moyenne glissante', 'Filtre passe-bas premier ordre', etc.)
    - param_base : str, paramètre de référence spécifique à la méthode de filtrage (par exemple, 'Fenêtre taille 20', 'Tau 0.1', etc.)
    - temps : ndarray (array numpy), vecteur de temps

    Retourne :
    - None
    """
    if 'taille' in param_base:
        params_filtrage = [param_base, f'{param_base} / 10', f'{param_base} * 10']
        method = 'moyenne_glissante'
    elif 'tau' in param_base:
        params_filtrage = [param_base, f'{param_base} / 10', f'{param_base} * 10']
        method = 'filtre_passe_bas_premier_ordre'
    else:
        params_filtrage = [param_base, f'{param_base} / 10', f'{param_base} * 10']
        method = 'filtre_passe_bas_second_ordre'

    if method == 'moyenne_glissante':
        signals_filtres = [signal_filtre_base,
                           moyenne_glissante(signal_brut, int(param_base.split()[-1]) // 10),
                           moyenne_glissante(signal_brut, int(param_base.split()[-1]) * 10)]
    elif method == 'filtre_passe_bas_premier_ordre':
        signals_filtres = [signal_filtre_base,
                           filtre_passe_bas_premier_ordre(signal_brut, float(param_base.split()[-1]) / 10),
                           filtre_passe_bas_premier_ordre(signal_brut, float(param_base.split()[-1]) * 10)]
    elif method == 'filtre_passe_bas_second_ordre':
        signals_filtres = [signal_filtre_base,
                           filtre_passe_bas_second_ordre(signal_brut, float(param_base.split()[-1]) / 10, xi),
                           filtre_passe_bas_second_ordre(signal_brut, float(param_base.split()[-1]) * 10, xi)]

    num_filtrage = len(signals_filtres)
    fig, axes = plt.subplots(num_filtrage, 1, figsize=(10, 6))

    for i in range(num_filtrage):
        ax = axes[i] if num_filtrage > 1 else axes  # Si une seule méthode de filtrage, utiliser le même axe
        ax.plot(temps, signal_brut, label='Signal brut')
        ax.plot(temps, signals_filtres[i], label=f'Signal filtré ({nom_filtrage} - {params_filtrage[i]})', linestyle='--')
        ax.set_xlabel('Temps (s)')
        ax.set_ylabel('Amplitude')
        ax.set_title(f'Comparaison du signal brut et du signal filtré ({nom_filtrage} - {params_filtrage[i]})')
        ax.legend()
        ax.grid(True)

    plt.tight_layout()
    plt.show()



## Import du signal bruité
# Chargement des tableaux numpy à partir des fichiers .npy
signal_EMG = np.load('signal_EMG.npy')
temps = np.load('temps.npy')

# Maintenant, signal_EMG contient le même tableau que celui qui a été enregistré
# par le script précédent.

### Traitement par filtrage
## Moyenne glissante
def moyenne_glissante(signal, fenetre):
    """
    Applique un filtre de moyenne glissante à un signal

    Paramètres :
    - signal : ndarray (array numpy), signal à filtrer
    - fenetre : int, taille de la fenêtre de la moyenne glissante

    Retourne :
    - signal_filtre : ndarray (array numpy), signal filtré par moyenne glissante
    """
    # création d'un tableau de zéros de même forme(shape) et type
    signal_filtre = np.zeros_like(signal)
    for i in range(len(signal)):
        # Attention aux effets de bord
        """ à compléter """
        signal_filtre[i] = np.mean(signal[debut_fenetre:fin_fenetre])
    return signal_filtre

# Paramètre de la moyenne glissante
taille_fenetre = 20

# Applique la moyenne glissante au signal bruité
signal_filtre_moyenne_glissante = moyenne_glissante(signal_EMG, taille_fenetre)


# Trace le signal bruité et les signaux filtrés par moyenne glissante
plot_comparison(signal_EMG, signal_filtre_moyenne_glissante, 'Moyenne glissante', f'Fenêtre taille {taille_fenetre}', temps)

# Trace le signal bruité et les signaux filtrés par moyenne glissante pour trois tailles de fenêtre
plot_comparison_subplots(signal_EMG, signal_filtre_moyenne_glissante, 'Moyenne glissante', f'Fenêtre taille {taille_fenetre}', temps)

## Filtre passe bas du premier ordre
def filtre_passe_bas_premier_ordre(signal, tau):
    """
    Applique un filtre passe-bas du premier ordre à un signal

    Paramètres :
    - signal : ndarray (array numpy), signal à filtrer
    - tau : float, constante de temps du premier ordre

    Retourne :
    - signal_filtre : ndarray (array numpy), signal filtré par le filtre passe-bas du premier ordre
    """
    signal_filtre = np.zeros_like(signal)

    """ à compléter """
    for i in range(1, len(signal)):
        signal_filtre[i] =    """ à compléter """

    return signal_filtre

# Paramètre du filtre passe-bas du premier ordre
tau = 0.1

# Appliquer le filtre passe-bas du premier ordre au signal bruité
signal_filtre_passe_bas_premier_ordre = filtre_passe_bas_premier_ordre(signal_EMG, tau)

# Tracer le signal bruité et le signal filtré par le filtre passe-bas du premier ordre
plot_comparison(signal_EMG, signal_filtre_passe_bas_premier_ordre, 'Passe-bas premier ordre', rf'Signal filtré ($\tau$ = {tau})', temps)

# Trace le signal bruité et les signaux filtrés par filtre passe-bas du premier ordre pour trois valeurs de tau
plot_comparison_subplots(signal_EMG, signal_filtre_passe_bas_premier_ordre, 'Passe-bas premier ordre', rf'Signal filtré $\tau$ = {tau}', temps)

## Filtre passe bas du second ordre
def filtre_passe_bas_second_ordre(signal, omega0, xi):
    """
    Applique un filtre passe-bas du second ordre à un signal

    Paramètres :
    - signal : ndarray (array numpy), signal à filtrer
    - omega0 : float, pulsation de coupure du filtre passe-bas du second ordre
    - xi : float, facteur d'amortissement du filtre passe-bas du second ordre

    Retourne :
    - signal_filtre : ndarray (array numpy), signal filtré par le filtre passe-bas du second ordre
    """


    # Initialisation du signal filtré
    signal_filtre = np.zeros_like(signal)
    """ à compléter """

    # Coefficients du filtre
    a =    """ à compléter """
    b =    """ à compléter """
    c =    """ à compléter """

    # Application du filtre
    for i in range(1, len(signal) - 1):
        signal_filtre[i + 1] =    """ à compléter """

    return signal_filtre

# Paramètres du filtre passe-bas du second ordre
omega0 = 10
xi = 1

# Appliquer le filtre passe-bas du second ordre au signal bruité
signal_filtre_passe_bas_second_ordre = filtre_passe_bas_second_ordre(signal_EMG, omega0, xi)

# Tracer le signal bruité et le signal filtré par le filtre passe-bas du second ordre
plot_comparison(signal_EMG, signal_filtre_passe_bas_second_ordre, 'Passe-bas second ordre', rf'Signal filtré ($\omega_0$ = {omega0}, $\xi$ = {xi})', temps)

# Tracer le signal bruité et le signal filtré par le filtre passe-bas du second ordre
plot_multiple_filter_comparison(signal_EMG, [signal_filtre_passe_bas_second_ordre, filtre_passe_bas_second_ordre(signal_EMG, omega0*2,xi),filtre_passe_bas_second_ordre(signal_EMG, omega0*0.1,xi)], ['Passe-bas second ordre','Passe-bas second ordre','Passe-bas second ordre'], [rf'$\omega_0$ = {omega0}, $\xi$ = {xi}', rf'$\omega_0$ = {omega0} * 5, $\xi$ = {xi}', rf'$\omega_0$ = {omega0} * 0.1, $\xi$ = {xi}'], temps)




