summaryrefslogtreecommitdiff
path: root/evaluate.py
diff options
context:
space:
mode:
Diffstat (limited to 'evaluate.py')
-rw-r--r--evaluate.py34
1 files changed, 22 insertions, 12 deletions
diff --git a/evaluate.py b/evaluate.py
index 7a39db4..7f996dc 100644
--- a/evaluate.py
+++ b/evaluate.py
@@ -76,10 +76,10 @@ def plot_training_curves(
76 ax1.set_title('Training and Validation Loss') 76 ax1.set_title('Training and Validation Loss')
77 ax1.legend() 77 ax1.legend()
78 78
79 ax2.plot(epochs, val_f1s, label='Val weighted F1', color='green') 79 ax2.plot(epochs, val_f1s, label='Val macro F1', color='green')
80 ax2.set_xlabel('Epoch') 80 ax2.set_xlabel('Epoch')
81 ax2.set_ylabel('Weighted F1') 81 ax2.set_ylabel('Macro F1')
82 ax2.set_title('Validation Weighted F1') 82 ax2.set_title('Validation Macro F1')
83 ax2.legend() 83 ax2.legend()
84 84
85 plt.tight_layout() 85 plt.tight_layout()
@@ -127,18 +127,28 @@ def generate_summary(all_results: dict, save_dir: str) -> None:
127 writer.writerows(rows) 127 writer.writerows(rows)
128 print(f"\nSummary saved to {csv_path}") 128 print(f"\nSummary saved to {csv_path}")
129 129
130 # Comparison bar chart 130 # Comparison bar chart — show both macro and weighted F1
131 labels = [f"{r['model']}\nCLAHE={r['clahe']}" for r in rows] 131 x_labels = [f"{r['model']}\nCLAHE={r['clahe']}" for r in rows]
132 w_f1_vals = [r['weighted_f1'] for r in rows] 132 w_f1_vals = [r['weighted_f1'] for r in rows]
133 133 m_f1_vals = [r['macro_f1'] for r in rows]
134 fig, ax = plt.subplots(figsize=(10, 5)) 134
135 bars = ax.bar(labels, w_f1_vals, color=['#4C72B0', '#DD8452'] * 3) 135 x = np.arange(len(x_labels))
136 ax.set_ylabel('Weighted F1-Score') 136 bar_width = 0.35
137 ax.set_title('Weighted F1 Comparison Across Experiments') 137 fig, ax = plt.subplots(figsize=(12, 5))
138 bars_w = ax.bar(x - bar_width / 2, w_f1_vals, bar_width, label='Weighted F1', color='#4C72B0')
139 bars_m = ax.bar(x + bar_width / 2, m_f1_vals, bar_width, label='Macro F1', color='#DD8452')
140 ax.set_ylabel('F1-Score')
141 ax.set_title('F1 Comparison Across Experiments')
142 ax.set_xticks(x)
143 ax.set_xticklabels(x_labels)
138 ax.set_ylim(0, 1.0) 144 ax.set_ylim(0, 1.0)
139 for bar, val in zip(bars, w_f1_vals): 145 ax.legend()
146 for bar, val in zip(bars_w, w_f1_vals):
147 ax.text(bar.get_x() + bar.get_width() / 2, bar.get_height() + 0.01,
148 f'{val:.3f}', ha='center', va='bottom', fontsize=8)
149 for bar, val in zip(bars_m, m_f1_vals):
140 ax.text(bar.get_x() + bar.get_width() / 2, bar.get_height() + 0.01, 150 ax.text(bar.get_x() + bar.get_width() / 2, bar.get_height() + 0.01,
141 f'{val:.4f}', ha='center', va='bottom', fontsize=9) 151 f'{val:.3f}', ha='center', va='bottom', fontsize=8)
142 plt.tight_layout() 152 plt.tight_layout()
143 chart_path = os.path.join(save_dir, 'comparison_chart.png') 153 chart_path = os.path.join(save_dir, 'comparison_chart.png')
144 fig.savefig(chart_path, dpi=150) 154 fig.savefig(chart_path, dpi=150)