Commit 6bfb5eec authored by Delvallez Delvallez's avatar Delvallez Delvallez

fullreplace pour l'automatisation des perturbations et les statistiques pour les composants

parent c5c42ada
......@@ -18,6 +18,8 @@
# * `mot+` : ajout de `mot` au début du document ;
# * `+mot` : ajout de `mot` à la fin du document.
# La perturbation spéciale `full_replace` correspond à l'application de toutes les perturbations `replace` fournies appliquées en une fois.
# Exemple :
# ```
......@@ -45,6 +47,7 @@
# * `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` ;
......@@ -96,15 +99,16 @@ text_field = "text" # default : "text"
is_irdata = False
query_id_subset=["1"] # default : None
perturbations_path = "perturbations_mainV2_elephantrose.txt"
perturbations_path = "res/mE5-MainV2/exp6/test/perturbation_polysemie_mainV2.txt"
full_replace = True
gen_matrices = True
gen_mean = True
gen_std = True
res_path = "res/mE5-MainV2/elephants-roses/"
res_path = "res/mE5-MainV2/exp6/test/"
load_values = True # default = False
load_path = "res/mE5-MainV2/elephants-roses/perts_values.json"
save_values = False # default = False
load_values = False # default = False
load_path = "res/mE5-MainV2/exp6/test/perts_values.json"
save_values = True # default = False
from mechir import Dot
......@@ -136,6 +140,14 @@ def param_prepend(doc, mot="microwave"):
def param_replace(doc, mot_orig="microwave", mot_rempl="toaster"):
return doc.replace(mot_orig, mot_rempl)
def param_full_replace(doc, replacmts):
res = doc
for replacemt in replacmts:
assert len(replacemt) == 3 , "Les perturbations de type full-replace doivent se composer de perturbations de type replace seulement"
res.replace(replacemt[1], replacemt[2])
return res
# génération des perturbations à produire à partir du document
def generer_transformations(fichier_regles):
......@@ -464,12 +476,13 @@ def calculate_components(param_pert_dot_dataloader, gen_mean, gen_std):
patching_head_outputs.append(patch_head_out)
print(f"Nombre 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 :
mean_head_outputs = torch.mean(torch.stack([tens.nan_to_num() for tens,_ in patching_head_outputs]), axis=0).detach().to("cpu").numpy()
mean_head_outputs = torch.mean(patching_head_outputs_cleaned, axis=0).detach().to("cpu").numpy()
if gen_std :
std_head_outputs = torch.std(torch.stack([tens.nan_to_num() for tens,_ in patching_head_outputs]), axis=0).detach().to("cpu").numpy()
std_head_outputs = torch.std(patching_head_outputs_cleaned, axis=0).detach().to("cpu").numpy()
return mean_head_outputs, std_head_outputs
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"):
if is_irdata:
......@@ -481,59 +494,85 @@ def load_dataset(data_path, is_irdata, query_id_subset = None, query_field = "te
# print("Number of query,doc pairs in dataset:", len(dataset))
return dataset
# Calcul ou extraction des valeurs
if load_values:
perts_values = load_pert_result(load_path)
# print(perts_values)
else:
# Récup du modèle
dot_model = Dot(dot_model_name)
# Recup dataset
dataset = load_dataset(data_path, is_irdata, query_id_subset, query_field, text_field)
perts, perts_name = generer_transformations(perturbations_path)
perts_values = {}
for i in range(len(perts)):
print(perts[i], "="*50)
if perts[i][0]=="replace":
pert_fun = perturbation(lambda texte : param_replace(texte, perts[i][1], perts[i][2]))
elif perts[i][0]=="prepend":
pert_fun = perturbation(lambda texte : param_prepend(texte, perts[i][1]))
elif perts[i][0]=="append":
pert_fun = perturbation(lambda texte : param_append(texte, perts[i][1]))
def one_pert_traitement(
pert, dot_model,
dataset,
full_replace = False,
gen_matrices = False,
gen_mean=False,
gen_std=False
):
if full_replace:
pert_fun = perturbation(lambda texte : param_full_replace(texte, pert))
pert_type = "replace"
elif pert[0]=="replace":
pert_fun = perturbation(lambda texte : param_replace(texte, pert[1], pert[2]))
pert_type = pert[0]
elif pert[0]=="prepend":
pert_fun = perturbation(lambda texte : param_prepend(texte, pert[1]))
pert_type = pert[0]
elif pert[0]=="append":
pert_fun = perturbation(lambda texte : param_append(texte, pert[1]))
pert_type = pert[0]
else:
raise RuntimeError("Unreachable")
param_pert_dot_collator = DotDataCollator(dot_model.tokenizer, pert_fun, perturb_type=perts[i][0])
param_pert_dot_collator = DotDataCollator(dot_model.tokenizer, pert_fun, perturb_type=pert_type)
param_pert_dot_dataloader = DataLoader(dataset, collate_fn=param_pert_dot_collator)
# param_pert_batch = next(iter(param_pert_dot_dataloader))
# print("PERTURBATION",i,":", perts_name[i])
# print("PERTURBATION",i,":", pert_name)
# pretty_print_triplets(param_pert_batch, dot_model.tokenizer, num=2)
baseline_perf = []
perturbed_perf = []
calculate_performance(dot_model, param_pert_dot_dataloader, baseline_perf, perturbed_perf, diags=True)
# plot_scores(baseline_perf, perturbed_perf, perts_name[i], save_path=res_path+"PerturbationScore_"+perts_name[i]+".png")
if gen_matrices:
mean_head_outputs, std_head_outputs = calculate_components(param_pert_dot_dataloader, gen_mean, gen_std)
# if mean_head_outputs is not None:
# plot_components_V2(mean_head_outputs, title="Components Patching Results (mean) for "+ perts_name[i], save_path=res_path+"ComponentPatching-mean_"+perts_name[i]+"V2.png")
# if std_head_outputs is not None:
# plot_components_V2(std_head_outputs, title="Components Patching Results (std) for "+ perts_name[i], save_path=res_path+"ComponentPatching-std_"+perts_name[i]+"V2.png")
mean_head_outputs, std_head_outputs, _ = calculate_components(param_pert_dot_dataloader, gen_mean, gen_std)
# Sauvegarde des valeurs pour la i^eme perturbation
perts_values[perts_name[i]] = {"baseline_perf": baseline_perf, "perturbed_perf": perturbed_perf}
# 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 :
perts_values[perts_name[i]]["patching_mean"] = mean_head_outputs
pert_values["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
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:
# Récup du modèle
dot_model = Dot(dot_model_name)
# Recup dataset
dataset = load_dataset(data_path, is_irdata, query_id_subset, query_field, text_field)
perts, perts_name = generer_transformations(perturbations_path)
perts_values = {}
if full_replace:
perts_values["full_replace"] = one_pert_traitement(
perts,
dot_model,
dataset,
full_replace=True,
gen_matrices= gen_matrices,
gen_mean= gen_mean,
gen_std = gen_std)
perts_values["full_replace"]["descr"] = perts_name
else:
for i in range(len(perts)):
print(perts[i], "="*50)
perts_values[perts_name[i]] = one_pert_traitement(
perts[i],
dot_model,
dataset,
full_replace=False,
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")
......
......@@ -93,6 +93,8 @@ def perturbation_type(pert):
return "append"
if pert[-1] =="+":
return "prepend"
if pert == "full_replace":
return "full_replace"
raise RuntimeWarning("Unreachable")
def stats(noeuds_perts, path):
......
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