Commit 9f84cee1 authored by Delvallez Delvallez's avatar Delvallez Delvallez

commentaires

parent 1b57123d
#!/usr/bin/env python3.12
# # Générateur de graphiques des performances pour des perturbations données
# # - Modèle : TAS-B
# # - Dataset : Vaswani part 1
# # - 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``
from mechir import Dot from mechir import Dot
from mechir.data import MechIRDataset, DotDataCollator from mechir.data import MechIRDataset, DotDataCollator
from mechir.perturb import perturbation from mechir.perturb import perturbation
...@@ -21,7 +32,7 @@ def param_prepend(doc, mot="microwave"): ...@@ -21,7 +32,7 @@ def param_prepend(doc, mot="microwave"):
def param_replace(doc, mot_orig="microwave", mot_rempl="toaster"): def param_replace(doc, mot_orig="microwave", mot_rempl="toaster"):
return doc.replace(mot_orig, mot_rempl) return doc.replace(mot_orig, mot_rempl)
# génération des perturbations à produire # génération des perturbations à produire à partir du document
def generer_transformations(fichier_regles): def generer_transformations(fichier_regles):
transformations = [] transformations = []
perts_name=[] perts_name=[]
...@@ -129,18 +140,18 @@ def plot_scores(baseline_scores, perturbed_scores, transformation_descr): ...@@ -129,18 +140,18 @@ def plot_scores(baseline_scores, perturbed_scores, transformation_descr):
def main():
# Récup du modèle
dot_model_name = "sebastian-hofstaetter/distilbert-dot-tas_b-b256-msmarco"
dot_model = Dot(dot_model_name)
# Récup du modèle # Recup dataset
dot_model_name = "sebastian-hofstaetter/distilbert-dot-tas_b-b256-msmarco" dataset = MechIRDataset("vaswani", query_id_subset=["1"])
dot_model = Dot(dot_model_name) print("Number of query,doc pairs in dataset:", len(dataset))
print("Query:", dataset._get_query("1"))
# Recup dataset perts, perts_name = generer_transformations("perturbations.txt")
dataset = MechIRDataset("vaswani", query_id_subset=["1"]) for i in range(len(perts)):
print("Number of query,doc pairs in dataset:", len(dataset))
print("Query:", dataset._get_query("1"))
perts, perts_name = generer_transformations("perturbations.txt")
for i in range(len(perts)):
if perts[i][0]=="replace": if perts[i][0]=="replace":
pert_fun = perturbation(lambda texte : param_replace(texte, perts[i][1], perts[i][2])) pert_fun = perturbation(lambda texte : param_replace(texte, perts[i][1], perts[i][2]))
elif perts[i][0]=="prepend": elif perts[i][0]=="prepend":
...@@ -160,3 +171,6 @@ for i in range(len(perts)): ...@@ -160,3 +171,6 @@ for i in range(len(perts)):
plot_scores(baseline_perf, perturbed_perf, perts_name[i]) plot_scores(baseline_perf, perturbed_perf, perts_name[i])
if __name__ == "__main__":
main()
\ No newline at end of file
#!/usr/bin/env python3 #!/usr/bin/env python3.12
# # À partir d'une liste de requêtes et d'une liste de documents (tous deux des strings) fourni dans deux csv
# # Extraction du vocabulaire présent dans les queries (df.vocab) et des statistiques suivantes pour chaque mot:
# # - mot : le mot
# # - idf : log10((nombre de documents + 1)/(nombre de documents contenant le mot+1))
# # - query_freq2 : nombre d'occurrence du mot dans les queries (2 occurrences dans une même query compte pour 2)
# # - docs_freq : nombre d'occurrence du mot dans les documents (2 occurrences dans un même documents compte pour 1)
"""
À partie d'une liste de requête et d'une liste de documents (tous deux des strings) fourni dans deux csv
Extraction du vocabulaire présent dans les queries (df.vocab) et des statistiques suivantes pour chaque mot:
- idf : log10((nombre de documents + 1)/(nombre de documents contenant le mot+1))
- query_freq2 : nombre d'occurrence du mot dans les queries (2 occurrences dans une même query compte pour 2)
- docs_freq : nombre d'occurrence du mot dans les documents (2 occurrences dans un même documents compte pour 1)
"""
import sys import sys
...@@ -83,7 +84,7 @@ def main(docs_path, query_path, csv_path, sort_criteria): ...@@ -83,7 +84,7 @@ def main(docs_path, query_path, csv_path, sort_criteria):
if __name__ == "__main__": if __name__ == "__main__":
if len(sys.argv) != 5: if len(sys.argv) != 5:
print("Usage : generation_idf.py <fichier_docs.csv> <fichier_query.csv> <csv_path.csv> <sort_criteria> \n where sort_criteria = 'query_freq2' or 'idf' or 'docs_freq' or 'alpha' or 'query_freq2/docs_freq'") print("Usage : generation_idf.py <fichier_docs.csv> <fichier_query.csv> <out_path.csv> <sort_criteria> \n where sort_criteria = 'query_freq2' or 'idf' or 'docs_freq' or 'alpha' or 'query_freq2/docs_freq'")
else: else:
main(sys.argv[1], sys.argv[2], sys.argv[3], sys.argv[4]) main(sys.argv[1], sys.argv[2], sys.argv[3], sys.argv[4])
\ No newline at end of file
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