import random import numpy as np from collections import deque import torch import torch.nn as nn import torch.optim as optim from safetensors.torch import load_file as load_safetensors from huggingface_hub import hf_hub_download # ------------------------------- # Neural Network Model # ------------------------------- class QNetwork(nn.Module): """ Deep Q-Network """ def __init__(self, state_dim, action_dim): super(QNetwork, self).__init__() self.model = nn.Sequential( nn.Linear(state_dim, 128), nn.ReLU(), nn.Linear(128, 128), nn.ReLU(), nn.Linear(128, action_dim) ) def forward(self, x): return self.model(x) # ------------------------------- # DQN Agent # ------------------------------- class RLAgent: """ Deep Q-Learning Agent for: Market Arbitrage + Grid-Aware Control """ def __init__( self, state_dim, action_dim=21, gamma=0.99, lr=1e-3, epsilon_start=1.0, epsilon_min=0.05, epsilon_decay=0.995, buffer_size=100000, batch_size=256, target_update_freq=500 ): self.state_dim = state_dim self.action_dim = action_dim self.gamma = gamma self.lr = lr self.epsilon = epsilon_start self.epsilon_min = epsilon_min self.epsilon_decay = epsilon_decay self.batch_size = batch_size self.target_update_freq = target_update_freq self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") # Discrete actions mapped to continuous power commands self.action_values = np.linspace(-1, 1, action_dim) # Experience Replay Buffer self.memory = deque(maxlen=buffer_size) # Networks self.policy_net = QNetwork(state_dim, action_dim).to(self.device) self.target_net = QNetwork(state_dim, action_dim).to(self.device) self.target_net.load_state_dict(self.policy_net.state_dict()) self.target_net.eval() self.optimizer = optim.Adam(self.policy_net.parameters(), lr=self.lr) self.loss_fn = nn.MSELoss() self.learn_step_counter = 0 # ------------------------------------------------ def act(self, state): """ Epsilon-greedy action selection. """ if np.random.rand() < self.epsilon: action_idx = random.randrange(self.action_dim) else: state = torch.FloatTensor(state).unsqueeze(0).to(self.device) with torch.no_grad(): q_values = self.policy_net(state) action_idx = torch.argmax(q_values).item() return self.action_values[action_idx] # ------------------------------------------------ def store(self, state, action, reward, next_state, done): """ Store experience in replay buffer. """ action_idx = np.argmin(np.abs(self.action_values - action)) self.memory.append((state, action_idx, reward, next_state, done)) # ------------------------------------------------ def learn(self): """ Sample mini-batch and perform learning step. """ if len(self.memory) < self.batch_size: return batch = random.sample(self.memory, self.batch_size) states, actions, rewards, next_states, dones = zip(*batch) states = torch.FloatTensor(states).to(self.device) actions = torch.LongTensor(actions).unsqueeze(1).to(self.device) rewards = torch.FloatTensor(rewards).unsqueeze(1).to(self.device) next_states = torch.FloatTensor(next_states).to(self.device) dones = torch.FloatTensor(dones).unsqueeze(1).to(self.device) # Current Q-values q_values = self.policy_net(states).gather(1, actions) # Target Q-values with torch.no_grad(): max_next_q = self.target_net(next_states).max(1)[0].unsqueeze(1) q_target = rewards + (1 - dones) * self.gamma * max_next_q loss = self.loss_fn(q_values, q_target) self.optimizer.zero_grad() loss.backward() self.optimizer.step() # Target network update self.learn_step_counter += 1 if self.learn_step_counter % self.target_update_freq == 0: self.target_net.load_state_dict(self.policy_net.state_dict()) # Epsilon decay if self.epsilon > self.epsilon_min: self.epsilon *= self.epsilon_decay # ------------------------------------------------ def save(self, path): torch.save(self.policy_net.state_dict(), path) # ------------------------------------------------ def load(self, path, use_safetensors=True): if "/" in path and not os.path.exists(path): # Probably a Hugging Face Repo ID print(f"📥 Downloading model from Hugging Face: {path}") repo_id = path filename = "dqn_energy_agent.safetensors" if use_safetensors else "dqn_energy_agent.pth" path = hf_hub_download(repo_id=repo_id, filename=filename) if path.endswith(".safetensors") or use_safetensors: state_dict = load_safetensors(path) else: state_dict = torch.load(path, map_location=self.device) self.policy_net.load_state_dict(state_dict) self.target_net.load_state_dict(self.policy_net.state_dict())