Skip to content

Commit cad2d7b

Browse files
committed
update tests
1 parent abc938f commit cad2d7b

6 files changed

Lines changed: 594 additions & 460 deletions

File tree

src/evaluation.py

Lines changed: 68 additions & 41 deletions
Original file line numberDiff line numberDiff line change
@@ -1,11 +1,13 @@
1-
import pandas as pd
2-
from sentence_transformers import SentenceTransformer, util
3-
import numpy as np
4-
from rapidfuzz import fuzz
1+
import random
2+
53
import Levenshtein
4+
import numpy as np
65
import open_clip
6+
import pandas as pd
77
import torch
8-
import random
8+
from rapidfuzz import fuzz
9+
from sentence_transformers import SentenceTransformer, util
10+
911

1012
def set_seed(seed=42):
1113
"""Set all random seeds for reproducibility"""
@@ -16,47 +18,55 @@ def set_seed(seed=42):
1618
torch.backends.cudnn.deterministic = True
1719
torch.backends.cudnn.benchmark = False
1820

19-
21+
2022
minilm_model = SentenceTransformer("sentence-transformers/all-MiniLM-L6-v2")
21-
biobert_model = SentenceTransformer("pritamdeka/BioBERT-mnli-snli-scinli-scitail-mednli-stsb")
23+
biobert_model = SentenceTransformer(
24+
"pritamdeka/BioBERT-mnli-snli-scinli-scitail-mednli-stsb"
25+
)
26+
2227

2328
# === Helpers ===
2429
def preprocess(text):
2530
return text.lower().strip()
2631

32+
2733
def split_and_clean(text):
2834
if pd.isna(text) or not text.strip():
2935
return []
3036
return [preprocess(x) for x in text.split(",") if x.strip()]
3137

38+
3239
def post_process(column):
33-
return (
34-
column.fillna(0.0)
35-
.round(4)
36-
.astype(float)
37-
.replace([np.inf, -np.inf], 0.0)
38-
)
40+
return column.fillna(0.0).round(4).astype(float).replace([np.inf, -np.inf], 0.0)
41+
3942

4043
# === Similarity Functions ===
41-
44+
45+
4246
def miniLM_pairwise_avg_similarity(list1, list2):
4347
if not list1 or not list2:
4448
return np.nan
4549
emb1 = minilm_model.encode(list1, convert_to_tensor=True)
4650
emb2 = minilm_model.encode(list2, convert_to_tensor=True)
4751
scores = util.cos_sim(emb1, emb2).cpu().numpy()
48-
semantic_scores = np.mean(np.max(scores, axis=1)) # average of best matches from list1 to list2
52+
semantic_scores = np.mean(
53+
np.max(scores, axis=1)
54+
) # average of best matches from list1 to list2
4955
return semantic_scores
5056

57+
5158
def biobert_pairwise_avg_similarity(list1, list2):
5259
if not list1 or not list2:
5360
return np.nan
5461
emb1 = biobert_model.encode(list1, convert_to_tensor=True)
5562
emb2 = biobert_model.encode(list2, convert_to_tensor=True)
5663
scores = util.cos_sim(emb1, emb2).cpu().numpy()
57-
semantic_scores = np.mean(np.max(scores, axis=1)) # average of best matches from list1 to list2
64+
semantic_scores = np.mean(
65+
np.max(scores, axis=1)
66+
) # average of best matches from list1 to list2
5867
return semantic_scores
5968

69+
6070
def fuzzy_pairwise_avg_similarity(list1, list2):
6171
# Fuzzy distance scores
6272
if not list1 or not list2:
@@ -68,10 +78,16 @@ def edit_pairwise_avg_similarity(list1, list2):
6878
# edit distance scores
6979
if not list1 or not list2:
7080
return np.nan
71-
return np.mean([
72-
max(1 - (Levenshtein.distance(s1, s2) / max(len(s1), len(s2))) for s2 in list2)
73-
for s1 in list1
74-
])
81+
return np.mean(
82+
[
83+
max(
84+
1 - (Levenshtein.distance(s1, s2) / max(len(s1), len(s2)))
85+
for s2 in list2
86+
)
87+
for s1 in list1
88+
]
89+
)
90+
7591

7692
# === Compute Similarity Columns ===
7793
def compute_similarity_column(sim_func, col_prefix, model_name, reader, df):
@@ -80,47 +96,61 @@ def compute_similarity_column(sim_func, col_prefix, model_name, reader, df):
8096
df[col_name] = df.apply(
8197
lambda row: sim_func(
8298
split_and_clean(row.get(reader, "")),
83-
split_and_clean(row.get(f"Protocol {model_name}", ""))
99+
split_and_clean(row.get(f"Protocol {model_name}", "")),
84100
),
85-
axis=1
101+
axis=1,
86102
)
87103
df[col_name] = post_process(df[col_name])
88104
return col_name
89105

90-
def evaluation_func(model_name, fine_tuning_method, reader):
106+
107+
def evaluation_func(model_name, fine_tuning_method, layers, reader):
91108
# === Config ===
92109
set_seed(42) # Set seed for reproducibility
93110
data_path = ""
94111

95-
if fine_tuning_method == 'none':
112+
if fine_tuning_method == "none":
96113
data_path = "/home/s33zganj/Machine-Learning-Based-Automation-of-MRI-Brain-Protocol-Selection/data/evaluation.xlsx"
97-
#if dataset_name == 'evaluation':
114+
# if dataset_name == 'evaluation':
98115
# data_path = "/home/s33zganj/Machine-Learning-Based-Automation-of-MRI-Brain-Protocol-Selection/data/evaluation.xlsx"
99-
#elif dataset_name == 'test':
116+
# elif dataset_name == 'test':
100117
# data_path = "/home/s33zganj/Machine-Learning-Based-Automation-of-MRI-Brain-Protocol-Selection/data/test.xlsx"
101-
118+
102119
else:
103-
data_path = f"/home/s33zganj/Machine-Learning-Based-Automation-of-MRI-Brain-Protocol-Selection/data/evaluation_{fine_tuning_method}_tuning.xlsx"
104-
#if dataset_name == 'evaluation':
120+
data_path = f"/home/s33zganj/Machine-Learning-Based-Automation-of-MRI-Brain-Protocol-Selection/data/evaluation_{fine_tuning_method}{layers}_tuning.xlsx"
121+
# if dataset_name == 'evaluation':
105122
# data_path = f"/home/s33zganj/Machine-Learning-Based-Automation-of-MRI-Brain-Protocol-Selection/data/evaluation_{fine_tuning_method}_tuning.xlsx"
106-
#elif dataset_name == 'test':
123+
# elif dataset_name == 'test':
107124
# data_path = f"/home/s33zganj/Machine-Learning-Based-Automation-of-MRI-Brain-Protocol-Selection/data/test_{fine_tuning_method}_tuning.xlsx"
108125

109-
110-
111-
#reader = "Protokoll"
126+
# reader = "Protokoll"
112127
# llm = "llama"
113128

114129
# === Load Data ===
115130
df = pd.read_excel(data_path)
116131

117-
118132
# Compute and post-process scores
119133
cols = []
120-
cols.append(compute_similarity_column(miniLM_pairwise_avg_similarity, "MiniLM", model_name, reader, df))
121-
cols.append(compute_similarity_column(biobert_pairwise_avg_similarity, "BioBert", model_name, reader, df))
122-
cols.append(compute_similarity_column(fuzzy_pairwise_avg_similarity, "fs", model_name, reader, df))
123-
cols.append(compute_similarity_column(edit_pairwise_avg_similarity, "ds", model_name, reader, df))
134+
cols.append(
135+
compute_similarity_column(
136+
miniLM_pairwise_avg_similarity, "MiniLM", model_name, reader, df
137+
)
138+
)
139+
cols.append(
140+
compute_similarity_column(
141+
biobert_pairwise_avg_similarity, "BioBert", model_name, reader, df
142+
)
143+
)
144+
cols.append(
145+
compute_similarity_column(
146+
fuzzy_pairwise_avg_similarity, "fs", model_name, reader, df
147+
)
148+
)
149+
cols.append(
150+
compute_similarity_column(
151+
edit_pairwise_avg_similarity, "ds", model_name, reader, df
152+
)
153+
)
124154

125155
# === Compute Final Scores ===
126156
print("\n=== Final Average Scores ===")
@@ -131,6 +161,3 @@ def evaluation_func(model_name, fine_tuning_method, reader):
131161
# === Save Results ===
132162
df.to_excel(data_path, index=False)
133163
print(f"\nSaved results to {data_path}")
134-
135-
136-

0 commit comments

Comments
 (0)