Skip to content

Commit 66d320f

Browse files
committed
feat(awa): add 3d performance metrics and enhanced visualization
1 parent 3ae4ca1 commit 66d320f

2 files changed

Lines changed: 104 additions & 36 deletions

File tree

scripts/experiment_awa.py

Lines changed: 39 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -185,6 +185,9 @@ def run_awa_experiment():
185185
ref_metrics_3d, _ = ecm.generate_pareto_front_3d(
186186
num_points_per_dim=15, **data_kwargs
187187
)
188+
ref_front_3d = np.column_stack(
189+
[ref_metrics_3d["Return"], ref_metrics_3d["Risk"], ref_metrics_3d["Diversification"]]
190+
)
188191

189192
print(" Evaluating MOEA/D (3D)...")
190193
moead_3d = MOEADOptimizer(problem=prob_3d)
@@ -195,6 +198,19 @@ def run_awa_experiment():
195198
**moead_kwargs,
196199
**data_kwargs,
197200
)
201+
front_moead_3d = np.column_stack(
202+
[
203+
metrics_moead_3d["Return"],
204+
metrics_moead_3d["Risk"],
205+
metrics_moead_3d["Diversification"],
206+
]
207+
)
208+
209+
m_dict_3d_1 = {
210+
"IGD": calculate_igd(front_moead_3d, ref_front_3d),
211+
"Spacing": calculate_spacing(front_moead_3d),
212+
"Spread": calculate_spread(front_moead_3d, ref_front_3d),
213+
}
198214

199215
print(" Evaluating MOEA/D-AWA (3D)...")
200216
awa_3d = MOEADAWAOptimizer(problem=prob_3d)
@@ -206,6 +222,26 @@ def run_awa_experiment():
206222
**moead_kwargs,
207223
**data_kwargs,
208224
)
225+
front_awa_3d = np.column_stack(
226+
[
227+
metrics_awa_3d["Return"],
228+
metrics_awa_3d["Risk"],
229+
metrics_awa_3d["Diversification"],
230+
]
231+
)
232+
233+
m_dict_3d_2 = {
234+
"IGD": calculate_igd(front_awa_3d, ref_front_3d),
235+
"Spacing": calculate_spacing(front_awa_3d),
236+
"Spread": calculate_spread(front_awa_3d, ref_front_3d),
237+
}
238+
239+
print(
240+
f" 3D MOEA/D - IGD: {m_dict_3d_1['IGD']:.6f}, Spacing: {m_dict_3d_1['Spacing']:.6f}, Spread: {m_dict_3d_1['Spread']:.6f}"
241+
)
242+
print(
243+
f" 3D MOEA/D-AWA - IGD: {m_dict_3d_2['IGD']:.6f}, Spacing: {m_dict_3d_2['Spacing']:.6f}, Spread: {m_dict_3d_2['Spread']:.6f}"
244+
)
209245

210246
from src.visualization import plot_two_variants_comparison_3d
211247

@@ -215,7 +251,9 @@ def run_awa_experiment():
215251
metrics_awa_3d,
216252
name1="Standard MOEA/D",
217253
name2="MOEA/D-AWA",
218-
save_path=os.path.join(run_folder, "exp4_3d_comparison.gif"),
254+
metrics_dict1=m_dict_3d_1,
255+
metrics_dict2=m_dict_3d_2,
256+
save_path=os.path.join(run_folder, "awa_comparison_3d.gif"),
219257
)
220258

221259
print(f"AWA Experiments finished. Results in {run_folder}")

src/visualization.py

Lines changed: 65 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -776,82 +776,112 @@ def plot_two_variants_comparison_3d(
776776
metrics2,
777777
name1="Standard MOEA/D",
778778
name2="Variant",
779+
metrics_dict1=None,
780+
metrics_dict2=None,
779781
save_path=None,
780782
):
781783
"""
782-
Compares two algorithm variants against ECM in two separate 3D subplots.
784+
Compares two algorithm variants against ECM in two separate 3D subplots,
785+
including performance metrics in the legend.
783786
"""
784-
fig = plt.figure(figsize=(18, 8))
787+
fig = plt.figure(figsize=(20, 9))
785788
ax1 = fig.add_subplot(121, projection="3d")
786789
ax2 = fig.add_subplot(122, projection="3d")
790+
791+
# Define color palette
792+
c_ref = "#333333"
793+
c1 = "#1f77b4"
794+
c2 = "#ff7f0e"
795+
796+
def format_label(name, m_dict):
797+
if not m_dict:
798+
return name
799+
m_str = "\n".join([f"{k}: {v:.4f}" for k, v in m_dict.items()])
800+
return f"{name}\n{m_str}"
787801

788802
# Subplot 1: Variant 1 vs ECM
789803
if ref_metrics_3d is not None:
790804
ax1.scatter(
791805
ref_metrics_3d["Return"],
792806
ref_metrics_3d["Risk"],
793807
ref_metrics_3d["Diversification"],
794-
c="black",
808+
c=c_ref,
795809
marker="x",
796-
s=20,
810+
s=25,
797811
label="Reference (ECM)",
798-
alpha=0.2,
812+
alpha=0.25,
813+
zorder=1,
799814
)
815+
816+
label1 = format_label(name1, metrics_dict1)
800817
ax1.scatter(
801818
metrics1["Return"],
802819
metrics1["Risk"],
803820
metrics1["Diversification"],
804821
marker="o",
805-
s=40,
806-
label=name1,
807-
alpha=0.6,
808-
color="#1f77b4",
822+
s=50,
823+
label=label1,
824+
alpha=0.7,
825+
color=c1,
826+
edgecolors="white",
827+
linewidth=0.5,
828+
zorder=2,
809829
)
810-
ax1.set_title(f"{name1} vs ECM")
811-
ax1.set_xlabel("Return")
812-
ax1.set_ylabel("Risk")
813-
ax1.set_zlabel("Diversification")
814-
ax1.legend()
830+
ax1.set_title(f"{name1} Performance (3D)", fontsize=14, pad=20)
831+
ax1.set_xlabel("Return", fontsize=11)
832+
ax1.set_ylabel("Risk", fontsize=11)
833+
ax1.set_zlabel("Diversification", fontsize=11)
834+
ax1.legend(loc="upper left", frameon=True, shadow=True, fontsize=10)
835+
ax1.view_init(elev=25, azim=45)
815836

816837
# Subplot 2: Variant 2 vs ECM
817838
if ref_metrics_3d is not None:
818839
ax2.scatter(
819840
ref_metrics_3d["Return"],
820841
ref_metrics_3d["Risk"],
821842
ref_metrics_3d["Diversification"],
822-
c="black",
843+
c=c_ref,
823844
marker="x",
824-
s=20,
845+
s=25,
825846
label="Reference (ECM)",
826-
alpha=0.2,
847+
alpha=0.25,
848+
zorder=1,
827849
)
850+
851+
label2 = format_label(name2, metrics_dict2)
828852
ax2.scatter(
829853
metrics2["Return"],
830854
metrics2["Risk"],
831855
metrics2["Diversification"],
832856
marker="o",
833-
s=40,
834-
label=name2,
835-
alpha=0.6,
836-
color="#ff7f0e",
857+
s=50,
858+
label=label2,
859+
alpha=0.7,
860+
color=c2,
861+
edgecolors="white",
862+
linewidth=0.5,
863+
zorder=2,
837864
)
838-
ax2.set_title(f"{name2} vs ECM")
839-
ax2.set_xlabel("Return")
840-
ax2.set_ylabel("Risk")
841-
ax2.set_zlabel("Diversification")
842-
ax2.legend()
843-
844-
# Animation function: Rotate both views
845-
def update(frame):
846-
ax1.view_init(elev=20, azim=frame)
847-
ax2.view_init(elev=20, azim=frame)
848-
return (fig,)
865+
ax2.set_title(f"{name2} Performance (3D)", fontsize=14, pad=20)
866+
ax2.set_xlabel("Return", fontsize=11)
867+
ax2.set_ylabel("Risk", fontsize=11)
868+
ax2.set_zlabel("Diversification", fontsize=11)
869+
ax2.legend(loc="upper left", frameon=True, shadow=True, fontsize=10)
870+
ax2.view_init(elev=25, azim=45)
849871

850-
ani = FuncAnimation(fig, update, frames=np.arange(0, 360, 4), interval=50)
872+
plt.suptitle("3D Pareto Front Quality Comparison (AWA Optimization)", fontsize=16, y=0.98)
873+
plt.tight_layout(rect=[0, 0.03, 1, 0.95])
851874

852875
if save_path:
853876
if save_path.endswith(".gif"):
854-
ani.save(save_path, writer="pillow", fps=15)
877+
def update(frame):
878+
ax1.view_init(elev=25, azim=frame)
879+
ax2.view_init(elev=25, azim=frame)
880+
return fig,
881+
882+
from matplotlib.animation import FuncAnimation
883+
ani = FuncAnimation(fig, update, frames=np.arange(0, 360, 5), interval=100)
884+
ani.save(save_path, writer="pillow", dpi=100)
855885
else:
856-
plt.savefig(save_path)
886+
plt.savefig(save_path, dpi=150, bbox_inches="tight")
857887
plt.close()

0 commit comments

Comments
 (0)