Commit 18f1639b authored by Delvallez Delvallez's avatar Delvallez Delvallez

amélioration de l'en-tête, affichage des graphiques, sauvegarde et chagement de données calculées

parent 26a6de01
#!/usr/bin/env python
# coding: utf-8
# ## Générateur de graphiques des performances pour des perturbations données
# Modèle : E5
# Dataset : Challenge evalLLM
# Exploite une liste de pertubation définies dans un fichier `perturbation.txt` sous trois formes:
# - `mot -> mot` pour un remplacement (éventuelement avec le mot vide pour supprimer)
# - `mot +` pour ajouter la requête puis le mot à gauche du document
# - `+ mot` pour ajouter le mot puis la requête à droite du document
#
# Les graphiques sont enregistrés dans `./res/<descripteur-de-la-perturbation>.png`
#
# # Générateur de graphiques de performance pour l'analyse de perturbations textuelles
# ## Description
# Ce script permet d'évaluer l'impact de perturbations lexicales sur les scores produits par un modèle de retrieval dense. À partir d'un ensemble de couples *(requête, document)* et d'une liste de perturbations définies par l'utilisateur, il génère différentes statistiques et représentations graphiques comparant les performances obtenues sur les documents originaux et leurs versions perturbées.
# Les perturbations sont appliquées uniquement aux documents. Les scores de similarité entre les requêtes et les documents sont ensuite recalculés afin d'étudier la sensibilité du modèle à ces modifications.
# ## Format des perturbations
# Les perturbations doivent être décrites dans un fichier texte contenant une règle par ligne. Trois types de transformations sont supportés :
# * `mot1->mot2` : remplacement de `mot1` par `mot2` ;
# * `mot->` : suppression de `mot` ;
# * `mot+` : ajout de `mot` au début du document ;
# * `+mot` : ajout de `mot` à la fin du document.
# Exemple :
# ```
# drone->aéronef
# militaire->
# important+
# +confidentiel
# ```
# ## Données d'entrée
# Le script nécessite :
# * un modèle de retrieval compatible avec MechIR ;
# * un jeu de données contenant des couples *(requête, document)* ;
# * un fichier décrivant les perturbations à appliquer.
# Les données peuvent provenir soit :
# * d'un fichier CSV ;
# * d'une collection compatible avec `ir_datasets`.
# ### Paramètres concernant les entrées
# * `dot_model_name` : nom du modèle dense utilisé pour le calcul des représentations dans la base HuggingFace;
# * `perturbations_path` : chemin vers le fichier de perturbations ;
# * `data_path` : chemin vers le fichier CSV ou identifiant de la collection dans la librairie `ir_datasets`;
# * `is_irdata` : indique si les données proviennent d'`ir_datasets` ;
# Dans le cas où le dataset provient d'un CSV:
# * `query_field` : nom de la colonne contenant les requêtes ;
# * `text_field` : nom de la colonne contenant les documents ;
# Dans le cas où le dataset est présent dans `ir_datasets`:
# * `query_id_subset` : liste optionnelle des identifiants de requêtes à conserver.
# ## Résultats produits
# Selon les options activées, le script peut générer :
# * les distributions des scores avant et après perturbation ;
# * les matrices d'impact des perturbations pour chaque composant du modèle sous la forme :
# * de moyennes des perturbations des paires ;
# * les écarts-types associés ;
# ### Paramètres concernant les sorties
# * `gen_matrices` : calcule les matrices d'impact des perturbations sur les composants ;
# * `gen_mean` : génère les graphiques des moyennes ;
# * `gen_std` : génère les graphiques des écarts-types;
# * `res_path` : répertoire où enregistrer les graphiques.
# ## Sauvegarde
# Un fichier JSON permettant de sauvegarder les scores calculés pour éviter leur recomputation lors d'exécutions ultérieures.
# Pour conserver les résultats calculés lors d'une execution, définir `save_values = True`.
# Pour pour exploiter des résultats préalablement sauvegardés, définir `load_values = True`.
# Le fichier JSON est enregistré/ lu est défini par `load_path`.
## paramètres du document
dot_model_name = "intfloat/multilingual-e5-small" # "sebastian-hofstaetter/distilbert-dot-tas_b-b256-msmarco"
data_path = "../challenge_data/mainV2/paires_mainV2.csv" # "vaswani" "../../../data/data_challenge/pairesV2_generees_1pourcent.csv"
query_field = "question" # "question" default : "text"
query_field = "question" # default : "text"
text_field = "text" # default : "text"
is_irdata = False
query_id_subset=["1"] # default : None
perturbations_path = "perturb_test.txt"
perturbations_path = "perturbations_mainV2_elephantrose.txt"
gen_matrices = True
gen_mean = True
gen_std = True
res_path = "res/mE5-MainV2/elephants-roses/"
load_values = False # default = False
load_path = "res/mE5-MainV2/elephant-rose/perts_values.json"
save_values = True # default = False
load_values = True # default = False
load_path = "res/mE5-MainV2/elephants-roses/perts_values.json"
save_values = False # default = False
res_path = "res/mE5-MainV2/elephant-rose/" #datetime.datetime.now().strftime("%Y-%m-%d-%H-%M")}
from mechir import Dot
from mechir.data import MechIRDataset, MechDataset, DotDataCollator
......@@ -138,7 +213,7 @@ def pretty_print_triplets(batch, tokenizer, num=1):
# Helper function to calculate and store performances
def calculate_performance(model, dataloader, baseline_performance, perturbed_performance, diags=False):
for i, batch in tqdm(enumerate(dataloader)):
for _, batch in tqdm(enumerate(dataloader)):
# Get the queries, documents, and perturbed documents from the batch
queries = batch["queries"]
documents = batch["documents"]
......@@ -220,21 +295,108 @@ def plot_score_dists_mult(all_baseline_scores, all_perturbed_scores, plot_type="
return
VIEW_STYLES = {
"base" : {
"cmap" :"RdBu_r",
"vmin" : -1,
"vmax" : 1,
},
"base_adapt" : {
"cmap" :"RdBu_r",
"vmin" : -1,
"vmax" : 1,
},
"base_extr" : {
"cmap" :"RdBu_r",
"vmin" : -50,
"vmax" : 50,
},
"base_abs" : {
"cmap" : "Reds",
"vmin" : 0,
"vmax" : 1,
},
"base_resc" : {
"cmap" : "RdBu_r",
"vmin" : -1,
"vmax" : 1,
},
"std" : {
"cmap" : "magma_r",
"vmin" : 0,
"vmax" : 1,
},
"std_adapt" : {
"cmap" : "magma_r",
"vmin" : 0,
"vmax" : 1,
},
"std_max" : {
"cmap" : "magma_r",
"vmin" : 0,
"vmax" : 50,
},
"std_resc" : {
"cmap" : "magma_r",
"vmin" : 0,
"vmax" : 1,
},
"norm" : {
"cmap" : "RdBu_r",
"vmin" : -1,
"vmax" : 1,
},
"std_norm" : {
"cmap" : "magma_r",
"vmin" : 0,
"vmax" : 1,
},
"abs" : {
"cmap" : "magma_r",
"vmin" : 0,
"vmax" : 50,
},
}
def plot_components_V2(
data, # shape: (num_layers, num_heads) or (num_layers, num_heads + 1) if include_mlp=True
save_path=None,
title="Component Patching Results",
include_mlp=False,
view_style = "base"
):
data = data.astype(float)
plt.figure(figsize=(10, 6))
printed_data = data
view_params = VIEW_STYLES[view_style].copy()
if "norm" in view_style:
# printed_data = printed_data / np.linalg.norm(data)
printed_data = (printed_data - printed_data.mean()) /printed_data.std()
if "abs" in view_style:
printed_data = printed_data.absolute()
if "extr" in view_style:
view_params["vmax"] = (abs(printed_data.max())+abs(printed_data.min())) /2
view_params["vmin"] = - view_params["vmax"]
if "max" in view_style:
view_params["vmax"] = (abs(printed_data.max())+abs(printed_data.min())) /2
view_params["vmin"] = 0
if "adapt" in view_style:
if (printed_data > 20).any() :
view_params["vmax"] = 50
else:
view_params["vmax"] = 1
view_params["vmin"] = (VIEW_STYLES[view_style]["vmin"] * view_params["vmax"]) # si vmin est déjà à 0 ça reste 0, si vmin est -1, on obtient -vmax
if "resc" in view_style:
printed_data = (np.sign(printed_data) / np.absolute(printed_data).max()) * np.absolute(printed_data)
ax = sns.heatmap(
np.abs(data),
cmap="margma_r",
vmin=0,
vmax=40,
printed_data,
cmap=view_params["cmap"],
vmin=view_params["vmin"],
vmax=view_params["vmax"],
xticklabels=True,
yticklabels=True,
fmt=".2f",
......@@ -270,25 +432,42 @@ def load_pert_result(filepath):
return perts_values
def save_pertubs_values(result_file, save_path):
save_values = {}
for pert_name in result_file:
save_values[pert_name] = {}
save_values[pert_name]["baseline_perf"] = result_file[pert_name]["baseline_perf"]
save_values[pert_name]["perturbed_perf"] = result_file[pert_name]["perturbed_perf"]
if "patching_mean" in result_file[pert_name] :
save_values[pert_name]["patching_mean"] = result_file[pert_name]["patching_mean"].tolist()
if "patching_std" in result_file[pert_name]:
save_values[pert_name]["patching_std"] = result_file[pert_name]["patching_std"].tolist()
with open(save_path, 'w', encoding="utf-8") as result_file:
json.dump(save_values, result_file, indent=2)
def calculate_components(param_pert_dot_dataloader, gen_mean, gen_std):
patching_head_outputs = []
nan_cpt = 0
for _, batch in tqdm(enumerate(param_pert_dot_dataloader)):
queries = batch["queries"]
documents = batch["documents"]
perturbed_documents = batch["perturbed_documents"]
patch_head_out = dot_model.patch(queries, documents, perturbed_documents, patch_type="head_all")
# Interruption du calcul si la matrice contient un NaN
if patch_head_out[0].isnan().any().item() :
print("Component Patching Result non calculé : contient au moins un NaN")
return None, None
patching_head_outputs.append(patch_head_out)
nan_cpt += patch_head_out[0].isnan().sum()
# # Interruption du calcul si la matrice contient un NaN
# if patch_head_out[0].isnan().any().item() :
# print("Component Patching Result non calculé : contient au moins un NaN")
# return None, None
patching_head_outputs.append(patch_head_out)
print(f"Nombre de NaN trouvés : {nan_cpt}")
mean_head_outputs, std_head_outputs = None, None
if gen_mean :
mean_head_outputs = torch.mean(torch.stack([tens for tens,_ in patching_head_outputs]), axis=0).detach().to("cpu").numpy()
mean_head_outputs = torch.mean(torch.stack([tens.nan_to_num() for tens,_ in patching_head_outputs]), axis=0).detach().to("cpu").numpy()
if gen_std :
std_head_outputs = torch.std(torch.stack([tens for tens,_ in patching_head_outputs]), axis=0).detach().to("cpu").numpy()
std_head_outputs = torch.std(torch.stack([tens.nan_to_num() for tens,_ in patching_head_outputs]), axis=0).detach().to("cpu").numpy()
return mean_head_outputs, std_head_outputs
......@@ -350,25 +529,25 @@ else:
# Sauvegarde des valeurs pour la i^eme perturbation
perts_values[perts_name[i]] = {"baseline": baseline_perf, "perturbed": perturbed_perf}
perts_values[perts_name[i]] = {"baseline_perf": baseline_perf, "perturbed_perf": perturbed_perf}
if gen_matrices and mean_head_outputs is not None :
perts_values[perts_name[i]]["patching_mean"] = mean_head_outputs.tolist()
perts_values[perts_name[i]]["patching_mean"] = mean_head_outputs
if gen_matrices and std_head_outputs is not None:
perts_values[perts_name[i]]["patching_std"] = std_head_outputs.tolist()
perts_values[perts_name[i]]["patching_std"] = std_head_outputs
if save_values:
with open(f"{res_path}/perts_values.json", 'w', encoding="utf-8") as result_file:
json.dump(perts_values, result_file, indent=2)
save_pertubs_values(perts_values, f"{res_path}/perts_values.json")
# Génération des graphiques
print("Génération des graphiques")
for pert_name, pert_values in perts_values.items():
plot_scores(pert_values["baseline"], pert_values["perturbed"], pert_name, save_path=res_path+"PerturbationScore_"+pert_name+".png")
plot_scores(pert_values["baseline_perf"], pert_values["perturbed_perf"], pert_name, save_path=res_path+"PerturbationScore_"+pert_name+".png")
if "patching_mean" in pert_values and pert_values["patching_mean"] is not None:
plot_components(pert_values["patching_mean"], title="Components Patching Results (mean) for "+ pert_name, save_path=res_path+"ComponentPatching-mean_"+pert_name+".png")
plot_components_V2(pert_values["patching_mean"], title="Components Patching Results (mean) for "+ pert_name, save_path=res_path+"ComponentPatching-mean_"+pert_name+"V2.png")
plot_components_V2(pert_values["patching_mean"], title="Components Patching Results (mean) for "+ pert_name, save_path=res_path+"ComponentPatching-mean_"+pert_name+".png", view_style="base_adapt")
plot_components_V2(pert_values["patching_mean"], title="Components Patching Results (mean) for "+ pert_name, save_path=res_path+"ComponentPatching-mean_"+pert_name+"V2.png", view_style="base_resc")
if "patching_std" in pert_values and pert_values["patching_std"] is not None:
plot_components(pert_values["patching_std"], title="Components Patching Results (std) for "+ pert_name, save_path=res_path+"ComponentPatching-std_"+pert_name+".png")
plot_components_V2(pert_values["patching_std"], title="Components Patching Results (std) for "+ pert_name, save_path=res_path+"ComponentPatching-std_"+pert_name+"V2.png")
plot_components_V2(pert_values["patching_std"], title="Components Patching Results (std) for "+ pert_name, save_path=res_path+"ComponentPatching-std_"+pert_name+".png", view_style="std_adapt")
plot_components_V2(pert_values["patching_std"], title="Components Patching Results (std) for "+ pert_name, save_path=res_path+"ComponentPatching-std_"+pert_name+"V2.png", view_style="std_resc")
Markdown is supported
0% or
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or to comment