PowerGrid / train /trainer.py
saiteja020's picture
Deployment fix: removed binary files from history
75c7554
Raw
History Blame Contribute Delete
5.73 kB
import os
import sys
import numpy as np
import time
import subprocess
import argparse
import json
# Ensure project root is on sys.path
_ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), '..'))
if _ROOT not in sys.path:
sys.path.insert(0, _ROOT)
# Canonical models directory (matches llm_evaluate.py and evaluate.py)
_MODELS_DIR = os.path.join(os.path.dirname(__file__), 'models')
from openenv.client import OpenEnvClient
from agent.config import AgentConfig
from agent.replay_buffer import ReplayBuffer
from agent.actor_critic import SAC_Agent
def start_server():
import requests
# Check if a server is already running (e.g. from Docker)
try:
r = requests.get("http://127.0.0.1:8000/api/health", timeout=1)
if r.status_code == 200:
print("Detected existing OpenEnv Server (likely Docker). Using it.")
return None
except:
pass
print("Starting local OpenEnv Server on port 8000...")
env_vars = os.environ.copy()
env_vars["PYTHONPATH"] = _ROOT
server_process = subprocess.Popen(
[sys.executable, "backend/main.py"], # Use our established main.py
env=env_vars,
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL
)
time.sleep(5) # Wait for server startup
return server_process
def train(args):
server_process = start_server()
try:
config = AgentConfig()
client = OpenEnvClient(base_url="http://127.0.0.1:8000")
agent = SAC_Agent(config)
load_path = args.load_model if args.load_model else os.path.join(_MODELS_DIR, f"best_model_{args.task}")
best_reward = -1000000.0
model_loaded_ok = False
# Architecture mismatch check
if os.path.exists(load_path + "_actor.pth"):
try:
print(f"Attempting to resume from: {load_path}")
# Note: SAC architecture is different from TDD-ND.
# This will likely fail if loading an old TDD-ND checkpoint.
agent.load(load_path)
print("Successfully loaded model weights.")
model_loaded_ok = True
except Exception as e:
print(f"Architecture Mismatch: {e}")
print("Old TDD-ND weights incompatible with SAC. Starting fresh.")
model_loaded_ok = False
# Load metadata only if model loaded successfully (avoids stale TDD-ND rewards)
task_meta_path = os.path.join(_MODELS_DIR, f"best_model_{args.task}_meta.json")
if model_loaded_ok and os.path.exists(task_meta_path):
with open(task_meta_path, "r") as f:
meta = json.load(f)
best_reward = meta.get("best_reward", -1000000.0)
agent.total_it = meta.get("total_it", 0)
print(f"Resumed best reward for task '{args.task}': {best_reward:.2f}")
replay_buffer = ReplayBuffer(config.state_dim, config.action_dim, config.buffer_size)
print(f"Starting SAC Agent Training on task: {args.task}")
os.makedirs(_MODELS_DIR, exist_ok=True)
total_steps = 0
for ep in range(args.episodes):
state = client.reset(seed=ep + int(time.time())%1000, task=args.task)
ep_reward = 0
ep_steps = 0
done = False
while not done:
# SAC uses stochastic sampling during training
if total_steps < config.exploration_steps and not os.path.exists(load_path + "_actor.pth"):
action = np.random.uniform(-config.max_action, config.max_action, size=config.action_dim)
else:
action = agent.select_action(np.array(state), evaluate=False)
next_state, reward, terminated, truncated, _ = client.step(action)
done = terminated or truncated
replay_buffer.add(state, action, next_state, reward, done)
state = next_state
ep_reward += reward
ep_steps += 1
total_steps += 1
if replay_buffer.size > config.batch_size:
agent.train(replay_buffer)
# Print metrics
alpha_val = agent.alpha.item() if hasattr(agent, 'alpha') else 0.0
print(f"Ep {ep+1}/{args.episodes} | Reward: {ep_reward:.2f} | Steps: {ep_steps} | Alpha: {alpha_val:.4f}")
# Save best model
if ep_reward > best_reward:
best_reward = ep_reward
save_path = os.path.join(_MODELS_DIR, f"best_model_{args.task}")
agent.save(save_path)
with open(os.path.join(_MODELS_DIR, f"best_model_{args.task}_meta.json"), "w") as f:
json.dump({"best_reward": best_reward, "total_it": agent.total_it, "algorithm": "SAC"}, f)
print(f" --> New best model saved! Reward: {best_reward:.2f}")
except Exception as outer_e:
print(f"Training Error: {outer_e}")
finally:
if server_process:
print("Shutting down local server...")
server_process.terminate()
else:
print("Training session finished.")
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--task", type=str, default="hard", choices=["easy", "medium", "hard"])
parser.add_argument("--episodes", type=int, default=50) # Increased default for SAC
parser.add_argument("--load-model", type=str, default=None)
args = parser.parse_args()
train(args)