Commit 7124622a authored by Delvallez Delvallez's avatar Delvallez Delvallez

Passage en CLI de GeneratedAutomatizedPerturbation.py + fichier de sortie pour les NaN

parent b258ef00
......@@ -47,7 +47,6 @@
# * `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 ;
# * `full_replace` : indique si les perturbations fournies doivent être appliquées comme une seule perturbation `full_replace`ou non (/!\ Dans ce cas, les perturbations décrites dans `perturbations_path` doivent être toutes de type replace.)
# * `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` ;
......@@ -72,7 +71,6 @@
# ### 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.
......@@ -98,11 +96,9 @@ query_field = "question" # default : "text"
text_field = "text" # default : "text"
is_irdata = False
query_id_subset=["1"] # default : None
sub_dataset_criteria = "pertinent" # default None
sub_dataset_criteria = None # default None "pertinent"
perturbations_path = "res/mE5-MainV2/test/new_pert_test.json"
full_replace = True
gen_matrices = True
gen_mean = False
gen_std = True
res_path = "res/mE5-MainV2/test/"
......@@ -130,6 +126,8 @@ import numpy as np
from tqdm import tqdm
import json
from copy import deepcopy
from pathlib import Path
import argparse
#Perturbations parametrées
......@@ -458,7 +456,26 @@ def save_pertubs_values(results_values, save_path):
with open(save_path, 'w', encoding="utf-8") as file:
json.dump(save_values, file, indent=2)
def calculate_components(param_pert_dot_dataloader, gen_mean, gen_std):
class NaNLog:
_instance = None
def __init__(self):
self.txt = ""
@staticmethod
def get():
if NaNLog._instance is None:
NaNLog._instance = NaNLog()
return NaNLog._instance
def log(self, info):
self.txt += info
def save(self, path):
with open(path, "a") as f:
f.write(self.txt)
def calculate_components(param_pert_dot_dataloader, dot_model, gen_mean, gen_std, pert_name):
patching_head_outputs = []
nan_cpt = 0
for _, batch in tqdm(enumerate(param_pert_dot_dataloader)):
......@@ -475,7 +492,7 @@ def calculate_components(param_pert_dot_dataloader, gen_mean, gen_std):
# 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}")
NaNLog.get().log(f"\n\n{pert_name}\nNombre de NaN trouvés : {nan_cpt}")
mean_head_outputs, std_head_outputs = None, None
patching_head_outputs_cleaned = torch.stack([tens.nan_to_num() for tens,_ in patching_head_outputs])
if gen_mean :
......@@ -485,14 +502,14 @@ def calculate_components(param_pert_dot_dataloader, gen_mean, gen_std):
return mean_head_outputs, std_head_outputs, patching_head_outputs_cleaned
def load_dataset(data_path, is_irdata, query_id_subset = None, query_field = "text", text_field ="text"):
def load_dataset(data_path, is_irdata, subdataset_condition=None, query_id_subset = None, query_field = "text", text_field ="text"):
if is_irdata:
dataset = MechIRDataset(data_path, query_id_subset=query_id_subset)
# print("Number of query,doc pairs in dataset:", len(dataset))
else:
data = pd.read_csv(data_path)
if sub_dataset_criteria is not None:
data = data.query(sub_dataset_criteria)
if subdataset_condition is not None:
data = data.query(subdataset_condition)
dataset = MechDataset(data, query_field=query_field, text_field=text_field)
# print("Number of query,doc pairs in dataset:", len(dataset))
return dataset
......@@ -501,7 +518,6 @@ def one_pert_traitement(
pert,
dot_model,
dataset,
gen_matrices = False,
gen_mean=False,
gen_std=False
):
......@@ -530,35 +546,47 @@ def one_pert_traitement(
perturbed_perf = []
calculate_performance(dot_model, param_pert_dot_dataloader, baseline_perf, perturbed_perf, diags=True)
if gen_matrices:
mean_head_outputs, std_head_outputs, _ = calculate_components(param_pert_dot_dataloader, gen_mean, gen_std)
if gen_mean or gen_std:
mean_head_outputs, std_head_outputs, _ = calculate_components(param_pert_dot_dataloader, dot_model, gen_mean, gen_std, pert["name"])
# Sauvegarde des valeurs pour la perturbation
pert_values = {"baseline_perf": baseline_perf, "perturbed_perf": perturbed_perf}
if gen_matrices and mean_head_outputs is not None :
if gen_mean and mean_head_outputs is not None :
pert_values["patching_mean"] = mean_head_outputs
if gen_matrices and std_head_outputs is not None:
if gen_std and std_head_outputs is not None:
pert_values["patching_std"] = std_head_outputs
return pert_values
# Calcul ou extraction des valeurs
if load_values:
perts_values = load_pert_result(load_path)
# print(perts_values)
else:
def generate_values(
model,
dataset_path,
dataset_origin,
perturbations_path,
query_field = "question",
text_field = "text",
subdataset_condition = None,
gen_mean = False,
gen_std = False,
output = None
):
# Récup du modèle
dot_model = Dot(dot_model_name)
dot_model = Dot(model)
# Recup dataset
dataset = load_dataset(data_path, is_irdata, query_id_subset, query_field, text_field)
dataset = load_dataset(
dataset_path,
dataset_origin == "irdatasets",
subdataset_condition,
query_id_subset,
query_field,
text_field
)
match perturbations_path.split(".")[-1].strip():
case "txt":
print("Appel à generer_transformations_txt")
perts = generer_transformations_txt(perturbations_path)
case "json":
print("Appel à generer_transformations_json")
perts = generer_transformations_json(perturbations_path)
case _:
raise RuntimeError("Unreachable")
......@@ -572,24 +600,141 @@ else:
pert,
dot_model,
dataset,
gen_matrices= gen_matrices,
gen_mean= gen_mean,
gen_std = gen_std)
})
if save_values:
save_pertubs_values(perts_values, f"{res_path}/perts_values.json")
out_path = (
Path(output)
if output is not None
else Path(perturbations_path).with_name(
"perts_values.json"
)
)
print(f"Enregistrement dans {out_path}")
save_pertubs_values(perts_values, out_path)
log_path = (
Path(output).with_name(
"NbNaN.txt"
)
if output is not None
else Path(perturbations_path).with_name(
"NbNaN.txt"
)
)
NaNLog.get().save(log_path)
# Génération des graphiques
print("Génération des graphiques")
for pert_values in perts_values:
plot_scores(pert_values["results"]["baseline_perf"], pert_values["results"]["perturbed_perf"], pert_values["pert"]["name"], save_path=res_path+"PerturbationScore_"+pert_values["pert"]["name"]+".png")
def generation_graphiques(input_path, output=None):
perts_values = load_pert_result(input_path)
res_path = (
Path(output)
if output is not None
else Path(perturbations_path).parent
)
for pert_values in perts_values:
plot_scores(pert_values["results"]["baseline_perf"], pert_values["results"]["perturbed_perf"], pert_values["pert"]["name"], save_path=res_path/("PerturbationScore_"+pert_values["pert"]["name"]+".png"))
if "patching_mean" in pert_values["results"] and pert_values["results"]["patching_mean"] is not None:
plot_components_V2(pert_values["results"]["patching_mean"], title="Components Patching Results (mean) for "+ pert_values["pert"]["name"], save_path=res_path+"ComponentPatching-mean_"+pert_values["pert"]["name"]+".png", view_style="base_adapt")
plot_components_V2(pert_values["results"]["patching_mean"], title="Components Patching Results (mean) for "+ pert_values["pert"]["name"], save_path=res_path+"ComponentPatching-mean_"+pert_values["pert"]["name"]+"V2.png", view_style="base_resc")
plot_components_V2(pert_values["results"]["patching_mean"], title="Components Patching Results (mean) for "+ pert_values["pert"]["name"], save_path=res_path/("ComponentPatching-mean_"+pert_values["pert"]["name"]+".png"), view_style="base_adapt")
plot_components_V2(pert_values["results"]["patching_mean"], title="Components Patching Results (mean) for "+ pert_values["pert"]["name"], save_path=res_path/("ComponentPatching-mean_"+pert_values["pert"]["name"]+"V2.png"), view_style="base_resc")
if "patching_std" in pert_values["results"] and pert_values["results"]["patching_std"] is not None:
plot_components_V2(pert_values["results"]["patching_std"], title="Components Patching Results (std) for "+ pert_values["pert"]["name"], save_path=res_path+"ComponentPatching-std_"+pert_values["pert"]["name"]+".png", view_style="std_adapt")
plot_components_V2(pert_values["results"]["patching_std"], title="Components Patching Results (std) for "+ pert_values["pert"]["name"], save_path=res_path+"ComponentPatching-std_"+pert_values["pert"]["name"]+"V2.png", view_style="std_resc")
plot_components_V2(pert_values["results"]["patching_std"], title="Components Patching Results (std) for "+ pert_values["pert"]["name"], save_path=res_path/("ComponentPatching-std_"+pert_values["pert"]["name"]+".png"), view_style="std_adapt")
plot_components_V2(pert_values["results"]["patching_std"], title="Components Patching Results (std) for "+ pert_values["pert"]["name"], save_path=res_path/("ComponentPatching-std_"+pert_values["pert"]["name"]+"V2.png"), view_style="std_resc")
if __name__ == "__main__":
parser = argparse.ArgumentParser()
subparser = parser.add_subparsers()
values_p = subparser.add_parser("values")
values_p.add_argument(
"model",
type=str,
help="Nom du modèle (chemin dans l'arborescence ou identifiant HuggingFace)"
)
values_p.add_argument(
"dataset",
type=str,
help="Nom du dataset dans irdataset ou chemin dans l'arborescence"
)
values_p.add_argument(
"dataset_origin",
choices = ["irdatasets", "file"],
help= "Origine du dataset. Indique comment obtenir le dataset."
)
values_p.add_argument(
"perturbations",
type=str,
help= "Chemin vers le fichier descriptif des perturbations"
)
values_p.add_argument(
"--query_field",
default= "question",
help="champ du fichier des paires correspondant aux questions"
)
values_p.add_argument(
"--docs_field",
default="text",
help= "Champ du fichier des paires correspondant aux documents"
)
values_p.add_argument(
"--subdataset",
default=None,
type=str,
help="Critère de sélection d'une sous-partie du dataset"
)
values_p.add_argument(
"--gen_mean",
action = "store_true",
help="Déclenche le calcul des matrices de moyennes de sensibilité"
)
values_p.add_argument(
"--gen_std",
action = "store_true",
help="Déclenche le calcul des matrices d'écarts-types de sensibilité"
)
values_p.add_argument(
"--output",
"-o",
default=None,
help="Spécifie le nom du fichier de sauvegarde des valeurs (Par défaut, le dossier des perturbations est utilisé et le fichier est nommé perts_values.json)"
)
values_p.set_defaults(cmd="values")
graphs_p = subparser.add_parser("graphs")
graphs_p.add_argument(
"input",
type=str,
help="Fichier contenant les données à mettre sous forme de graphique"
)
graphs_p.add_argument(
"--output",
'-o',
help="Dossier de sauvegarde des graphiques (Par défaut, le dossier de l'input)"
)
graphs_p.set_defaults(cmd="graphs")
args = parser.parse_args()
if args.cmd == "values":
generate_values(
model= args.model,
dataset_path= args.dataset,
dataset_origin= args.dataset_origin,
perturbations_path= args.perturbations,
query_field= args.query_field,
text_field= args.docs_field,
subdataset_condition= args.subdataset,
gen_mean= args.gen_mean,
gen_std= args.gen_std,
output= args.output
)
elif args.cmd == "graphs":
generation_graphiques(
input_path= args.input,
output = args.output
)
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