Energizer / scripts /evaluation.py
Srirama-Mithilesh
Initialize final clean source-only repository with embedded weights
d7adedb
Raw
History Blame Contribute Delete
4.49 kB
import sys
import os
import numpy as np
import torch
import matplotlib.pyplot as plt
# Add root directory to sys.path for absolute imports
sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), '..')))
from openenv.grid_twin import GridTwinEnv
from openenv.models import Action
import glob
from agent.loader import load_trained_model
# =========================
# SETUP
# =========================
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
PLOT_DIR = "visualization"
os.makedirs(PLOT_DIR, exist_ok=True)
# 🧹 Clear existing plots to prevent accumulation
for f in glob.glob(f"{PLOT_DIR}/*.png"):
try:
os.remove(f)
except OSError:
pass
# =========================
# LOAD SAC MODEL
# =========================
actor = load_trained_model()
# =========================
# EVALUATION (7 DAYS)
# =========================
print("\n=== SAC RL Evaluation (Grid-Aware) ===")
total_rewards = []
total_profits = []
total_EA_all = 0
total_FR_all = 0
total_PS_all = 0
total_deg_all = 0
total_v_viols = 0
# Store Day 0 Trajectory
soc_list, power_list, demand_list, lmp_list = [], [], [], []
regd_list, voltage_list, line_load_list = [], [], []
num_days = 7
for day in range(num_days):
env = GridTwinEnv(seed=day)
obs = env.reset(episode=999).to_array() # Ensure High-Fidelity
done = False
total_reward = 0
while not done:
state = torch.tensor(obs, dtype=torch.float32).to(device).unsqueeze(0)
with torch.no_grad():
action = actor.get_action(state, deterministic=True)
next_obs, reward, done, info = env.step(Action(power=action))
total_reward += reward
obs = next_obs.to_array()
if day == 0:
soc_list.append(env.soc)
power_list.append(action)
demand_list.append(env.demand_series[env.t-1])
lmp_list.append(env.lmp_series[env.t-1])
regd_list.append(env.regd_series[env.t-1])
voltage_list.append(env.v_pu)
line_load_list.append(env.line_loading)
if env.v_pu < 0.95 or env.v_pu > 1.05:
total_v_viols += 1
day_profit = env.total_EA + env.total_FR
print(f"Day {day}: Reward={total_reward:7.2f} | Profit=${day_profit:7.2f} | EA=${env.total_EA:7.2f}")
total_rewards.append(total_reward)
total_profits.append(day_profit)
total_EA_all += env.total_EA
total_FR_all += env.total_FR
total_PS_all += env.total_PS
total_deg_all += env.total_deg
# =========================
# PERFORMANCE SUMMARY
# =========================
print("\n=== AGGREGATE STATS ===")
print(f"Avg Reward: {np.mean(total_rewards):.3f}")
print(f"Avg Daily Profit: ${np.mean(total_profits):.2f}")
print(f"Total Voltage Violations: {total_v_viols} (over {num_days*288} steps)")
print(f"Avg Battery Deg: {total_deg_all / num_days:.3f} MWh/day")
print("\n=== REVENUE BREAKDOWN (Avg/Day) ===")
print(f"Energy Arbitrage (EA): ${total_EA_all / num_days:7.2f}")
print(f"Freq Regulation (FR): ${total_FR_all / num_days:7.2f}")
print(f"Peak Stress Level: {total_PS_all / num_days:7.2f}")
# =========================
# PLOTS (ONLY DAY 0)
# =========================
t = np.arange(len(soc_list))
plt.rcParams['figure.figsize'] = [10, 5]
# Grid Physics (Consolidated)
plt.figure()
ax1 = plt.gca()
ax1.plot(t, voltage_list, color='teal', label="Voltage (pu)")
ax1.axhline(1.05, color='r', ls='--')
ax1.axhline(0.95, color='r', ls='--')
ax1.set_ylabel("Voltage (pu)")
ax2 = ax1.twinx()
ax2.plot(t, line_load_list, color='orange', alpha=0.5, label="Line Loading (%)")
ax2.axhline(80, color='brown', ls=':')
ax2.set_ylabel("Loading (%)")
plt.title("Day 0: Grid Constraints Compliance")
plt.savefig(f"{PLOT_DIR}/grid_compliance.png")
plt.close()
# Revenue Breakdown Pie
labels = ["Energy Arbitrage", "Freq Regulation"]
values = [max(0.1, total_EA_all), max(0.1, total_FR_all)]
plt.figure()
plt.pie(values, labels=labels, autopct='%1.1f%%', colors=['#4CAF50', '#2196F3'])
plt.title("Revenue Mix (7 Days)")
plt.savefig(f"{PLOT_DIR}/breakdown.png")
plt.close()
# SOC & Price
fig, ax1 = plt.subplots()
ax1.plot(t, lmp_list, color='gray', alpha=0.3, label="LMP")
ax1.set_ylabel("Price")
ax2 = ax1.twinx()
ax2.plot(t, soc_list, color='blue', label="SOC")
ax2.set_ylabel("SOC [0-1]")
plt.title("Day 0: Cycling Behavior vs Price")
plt.savefig(f"{PLOT_DIR}/soc_price.png")
plt.close()
print(f"\nEvaluation complete. Trajectories saved to {PLOT_DIR}/")