diff --git a/src/Car-Racing/doubledqn_model/double_dqn.py b/src/Car-Racing/doubledqn_model/double_dqn.py new file mode 100644 index 0000000000..6a4613b570 --- /dev/null +++ b/src/Car-Racing/doubledqn_model/double_dqn.py @@ -0,0 +1,225 @@ +import os +import sys +import torch +import numpy as np +from torch import nn +from torchrl.data import TensorDictReplayBuffer, LazyMemmapStorage +from tensordict import TensorDict +import datetime +import csv +import matplotlib.pyplot as plt # Import for plotting +import yaml +from pathlib import Path + +sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), '..'))) +from dqn_model.DQN_model import Agent as DQNAgent, SkipFrame # Import SkipFrame from DQN_model + +# Define is_ipython locally +import matplotlib +is_ipython = 'inline' in matplotlib.get_backend() +if is_ipython: + from IPython import display + +class DoubleDQNAgent(DQNAgent): + def __init__(self, state_space_shape, action_n, config_path='configs/double_dqn.yaml', load_state=False, load_model=None, **kwargs): + # Convert to Path object if it's a string + config_path = Path(config_path) if isinstance(config_path, str) else config_path + + # Load base config first + base_config_path = config_path.parent / 'dqn.yaml' + with open(base_config_path) as f: + base_config = yaml.safe_load(f) + + # Load and merge DoubleDQN specific config + with open(config_path) as f: + double_dqn_config = yaml.safe_load(f) + + # Merge configs with DoubleDQN having priority + self.config = {'hyperparameters': {}} + if 'hyperparameters' in base_config: + self.config['hyperparameters'].update(base_config['hyperparameters']) + if 'hyperparameters' in double_dqn_config: + self.config['hyperparameters'].update(double_dqn_config['hyperparameters']) + + # Set default values for DoubleDQN specific parameters if not provided + self.config['hyperparameters'].setdefault('tau', 0.005) + self.config['hyperparameters'].setdefault('update_target_every', 10000) + + # Initialize parent with explicit parameters + super().__init__( + state_space_shape=state_space_shape, + action_n=action_n, + config=self.config, + **kwargs + ) + + # Double DQN specific parameters + self.tau = self.config['hyperparameters']['tau'] + self.update_target_every = self.config['hyperparameters']['update_target_every'] + + # Initialize Double DQN networks + self._initialize_networks(load_state, load_model) + + def _initialize_networks(self, load_state, load_model): + """Initialize policy_net and target_net networks""" + self.policy_net = self._build_network().float().to(self.device) + self.target_net = self._build_network().float().to(self.device) + self.target_net.load_state_dict(self.policy_net.state_dict()) + self.optimizer = torch.optim.Adam(self.policy_net.parameters(), lr=self.hyperparameters['lr']) + # If loading state, do it after networks are initialized + if load_state and load_model: + self.load(os.path.join(self.save_dir, load_model)) + + def update_net(self, batch_size): + """Override update method for Double DQN logic""" + self.n_updates += 1 + states, actions, rewards, new_states, terminateds = self.get_samples(batch_size) + + # Current Q values + current_q = self.policy_net(states).gather(1, actions.unsqueeze(1)) + + # Double DQN target calculation + with torch.no_grad(): + next_actions = self.policy_net(new_states).argmax(1, keepdim=True) + next_q = self.target_net(new_states).gather(1, next_actions) + target_q = rewards.unsqueeze(1) + (1 - terminateds.float().unsqueeze(1)) * self.gamma * next_q + + loss = self.loss_fn(current_q, target_q) + self.optimizer.zero_grad() + loss.backward() + self.optimizer.step() + + # Soft update target network + if self.n_updates % self.update_target_every != 0: + for target_param, policy_param in zip(self.target_net.parameters(), self.policy_net.parameters()): + target_param.data.copy_(self.tau * policy_param.data + (1.0 - self.tau) * target_param.data) + + return current_q.mean().item(), loss.item() + + def _build_network(self): + """Same architecture as your DQN for compatibility""" + return nn.Sequential( + nn.Conv2d(self.state_shape[0], 16, kernel_size=8, stride=4), + nn.ReLU(), + nn.Conv2d(16, 32, kernel_size=4, stride=2), + nn.ReLU(), + nn.Flatten(), + nn.Linear(2592, 256), + nn.ReLU(), + nn.Linear(256, self.action_n) + ) + + def store(self, state, action, reward, new_state, terminated): + """Identical storage method to maintain compatibility""" + self.buffer.add(TensorDict({ + "state": torch.tensor(state), + "action": torch.tensor(action), + "reward": torch.tensor(reward), + "new_state": torch.tensor(new_state), + "terminated": torch.tensor(terminated) + }, batch_size=[])) + + def get_samples(self, batch_size): + """Identical sampling method""" + batch = self.buffer.sample(batch_size) + states = batch.get('state').float().to(self.device) + new_states = batch.get('new_state').float().to(self.device) + actions = batch.get('action').squeeze().to(self.device) + rewards = batch.get('reward').squeeze().to(self.device) + terminateds = batch.get('terminated').squeeze().to(self.device) + return states, actions, rewards, new_states, terminateds + + def take_action(self, state): + """Identical action selection""" + if np.random.rand() < self.epsilon: + action_idx = np.random.randint(self.action_n) + else: + state = torch.tensor(state, dtype=torch.float32, device=self.device).unsqueeze(0) + action_values = self.policy_net(state) + action_idx = torch.argmax(action_values, axis=1).item() + + # Decay epsilon + if self.epsilon > self.epsilon_min: + self.epsilon *= self.epsilon_decay + else: + self.epsilon = self.epsilon_min + + self.act_taken += 1 + return action_idx + + def save(self, save_dir, filename): + """Override save to use parent-compatible format""" + os.makedirs(save_dir, exist_ok=True) + model_path = os.path.join(save_dir, f"{filename}.pt") + + torch.save({ + 'updating_net_state_dict': self.policy_net.state_dict(), + 'frozen_net_state_dict': self.target_net.state_dict(), + 'optimizer_state_dict': self.optimizer.state_dict(), + 'epsilon': self.epsilon, + 'n_updates': self.n_updates, + 'config': self.config # Include config for consistency + }, model_path) + + print(f"Model weights saved to: {model_path}") + + def load(self, path): + """Override load to use parent-compatible format""" + checkpoint = torch.load(path, map_location=torch.device('cpu')) + self.policy_net.load_state_dict(checkpoint['updating_net_state_dict']) + self.target_net.load_state_dict(checkpoint['frozen_net_state_dict']) + self.optimizer.load_state_dict(checkpoint['optimizer_state_dict']) + self.epsilon = checkpoint['epsilon'] + self.n_updates = checkpoint['n_updates'] + + # Handle config mismatch if present + if 'config' in checkpoint: + for key in ['gamma', 'epsilon_decay', 'epsilon_min', 'tau']: + if checkpoint['config'].get(key) != getattr(self, key): + print(f"Warning: Config mismatch for {key} - " + f"Model: {checkpoint['config'].get(key)}, " + f"Current: {getattr(self, key)}") + + print(f"Loaded weights from {path}") + + def write_log(self, date_list, time_list, reward_list, length_list, loss_list, epsilon_list, log_filename='double_dqn_log.csv'): + """Identical logging method""" + if not os.path.exists(self.log_dir): + os.makedirs(self.log_dir) + rows = [ + ['date'] + date_list, + ['time'] + time_list, + ['reward'] + reward_list, + ['length'] + length_list, + ['loss'] + loss_list, + ['epsilon'] + epsilon_list + ] + with open(os.path.join(self.log_dir, log_filename), 'w') as csvfile: + csvwriter = csv.writer(csvfile) + csvwriter.writerows(rows) + +def plot_reward(episode_num, reward_list, n_steps): + plt.figure(1) + rewards_tensor = torch.tensor(reward_list, dtype=torch.float) + if len(rewards_tensor) >= 11: + eval_reward = torch.clone(rewards_tensor[-10:]) + mean_eval_reward = round(torch.mean(eval_reward).item(), 2) + std_eval_reward = round(torch.std(eval_reward).item(), 2) + plt.clf() + plt.title(f'Episode #{episode_num}: {n_steps} steps, ' + f'reward {mean_eval_reward}±{std_eval_reward}') + else: + plt.clf() + plt.title('Training...') + plt.xlabel('Episode') + plt.ylabel('Reward') + plt.plot(rewards_tensor.numpy()) + if len(rewards_tensor) >= 50: + reward_f = torch.clone(rewards_tensor[:50]) + means = rewards_tensor.unfold(0, 50, 1).mean(1).view(-1) + means = torch.cat((torch.ones(49) * torch.mean(reward_f), means)) + plt.plot(means.numpy()) + plt.pause(0.001) + if is_ipython: + display.display(plt.gcf()) + display.clear_output(wait=True) \ No newline at end of file diff --git a/src/Car-Racing/doubledqn_model/training_double_dqn.py b/src/Car-Racing/doubledqn_model/training_double_dqn.py new file mode 100644 index 0000000000..2d6b36b659 --- /dev/null +++ b/src/Car-Racing/doubledqn_model/training_double_dqn.py @@ -0,0 +1,231 @@ +import os +import sys +import matplotlib +import torch +import datetime +import csv +from pathlib import Path + +import gymnasium as gym +import gymnasium.wrappers as gym_wrap +import matplotlib.pyplot as plt +import numpy as np + +from double_dqn import plot_reward +from double_dqn import DoubleDQNAgent + +# Adjust path to import from parent directory +sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), '..'))) +from double_dqn_model.double_dqn import DoubleDQNAgent +from dqn_model.DQN_model import SkipFrame + +from double_dqn_model.double_dqn import plot_reward + +from gymnasium.spaces import Box +from tensordict import TensorDict +from torch import nn +from torchrl.data import TensorDictReplayBuffer, LazyMemmapStorage + +is_ipython = 'inline' in matplotlib.get_backend() +if is_ipython: + from IPython import display + +plt.ion() + +# Environment setup +env = gym.make("CarRacing-v3", continuous=False) +env = SkipFrame(env, skip=4) +from gymnasium.wrappers import GrayscaleObservation, ResizeObservation, FrameStackObservation +env = GrayscaleObservation(env) +env = ResizeObservation(env, (84, 84)) +env = FrameStackObservation(env, stack_size=4) + +# Initialize environment and get state shape +state, info = env.reset() +action_n = env.action_space.n +print(f"Action space size: {action_n}") +print(f"Action space details: {env.action_space}") + +# Ensure correct config path +config_path = Path(__file__).parent.parent / 'configs' / 'double_dqn.yaml' + +# Agent initialization +driver = DoubleDQNAgent( + state_space_shape=state.shape, + action_n=action_n, + config_path=config_path, + load_state=False, # Explicitly set to False if not loading a model + load_model=None # Set to model filename if loading +) + +# Verify config loaded +print(f"Using tau: {driver.tau}") +print(f"Target update every: {driver.update_target_every} steps") + +# Training parameters +batch_n = 32 +play_n_episodes = 2000 # Keep low for quick testing, adjust as needed +episode_epsilon_list = [] +episode_reward_list = [] +episode_length_list = [] +episode_loss_list = [] +episode_date_list = [] +episode_time_list = [] +episode = 0 +timestep_n = 2 +when2learn = 4 # in timesteps +when2sync = 5000 # in timesteps +when2save = 100000 # in timesteps +when2report = 5000 # in timesteps +when2eval = 50000 # in timesteps +when2log = 10 # in episodes +report_type = 'plot' # 'text', 'plot', None + +# Training loop +while episode <= play_n_episodes: + episode += 1 + episode_reward = 0 + episode_length = 0 + updating = True + loss_list = [] + episode_epsilon_list.append(driver.epsilon) + + while updating: + timestep_n += 1 + episode_length += 1 + + action = driver.take_action(state) + new_state, reward, terminated, truncated, info = env.step(action) + episode_reward += reward + driver.store(state, action, reward, new_state, terminated) + state = new_state + updating = not (terminated or truncated) + + if timestep_n % when2sync == 0: + # Sync target network manually (Double DQN handles this internally, but we can enforce it) + driver.target_net.load_state_dict(driver.policy_net.state_dict()) + + if timestep_n % when2save == 0: + driver.save(driver.save_dir, 'DoubleDQN') + + if timestep_n % when2learn == 0 and len(driver.buffer) >= batch_n: + q, loss = driver.update_net(batch_n) + loss_list.append(loss) + + if timestep_n % when2report == 0 and report_type == 'text': + print(f'Report: {timestep_n} timestep') + print(f' episodes: {episode}') + print(f' n_updates: {driver.n_updates}') + print(f' epsilon: {driver.epsilon}') + + if timestep_n % when2eval == 0 and report_type == 'text': + rewards_tensor = torch.tensor(episode_reward_list, dtype=torch.float) + eval_reward = torch.clone(rewards_tensor[-50:]) + mean_eval_reward = round(torch.mean(eval_reward).item(), 2) + std_eval_reward = round(torch.std(eval_reward).item(), 2) + + lengths_tensor = torch.tensor(episode_length_list, dtype=torch.float) + eval_length = torch.clone(lengths_tensor[-50:]) + mean_eval_length = round(torch.mean(eval_length).item(), 2) + std_eval_length = round(torch.std(eval_length).item(), 2) + + print(f'Evaluation: {timestep_n} timestep') + print(f' reward {mean_eval_reward}±{std_eval_reward}') + print(f' episode length {mean_eval_length}±{std_eval_length}') + print(f' episodes: {episode}') + print(f' n_updates: {driver.n_updates}') + print(f' epsilon: {driver.epsilon}') + + state, info = env.reset() + + episode_reward_list.append(episode_reward) + episode_length_list.append(episode_length) + episode_loss_list.append(np.mean(loss_list) if loss_list else 0) + now_time = datetime.datetime.now() + episode_date_list.append(now_time.date().strftime('%Y-%m-%d')) + episode_time_list.append(now_time.time().strftime('%H:%M:%S')) + + if report_type == 'plot': + draw_check = plot_reward(episode, episode_reward_list, timestep_n) + + if episode % when2log == 0: + driver.write_log( + episode_date_list, + episode_time_list, + episode_reward_list, + episode_length_list, + episode_loss_list, + episode_epsilon_list, + log_filename='DoubleDQN_log_test.csv' + ) + +if report_type == 'text': + rewards_tensor = torch.tensor(episode_reward_list, dtype=torch.float) + eval_reward = torch.clone(rewards_tensor[-100:]) + mean_eval_reward = round(torch.mean(eval_reward).item(), 2) + std_eval_reward = round(torch.std(eval_reward).item(), 2) + + lengths_tensor = torch.tensor(episode_length_list, dtype=torch.float) + eval_length = torch.clone(lengths_tensor[-100:]) + mean_eval_length = round(torch.mean(eval_length).item(), 2) + std_eval_length = round(torch.std(eval_length).item(), 2) + + print(f'Final evaluation: {timestep_n} timestep') + print(f' reward {mean_eval_reward}±{std_eval_reward}') + print(f' episode length {mean_eval_length}±{std_eval_length}') + print(f' episodes: {episode}') + print(f' n_updates: {driver.n_updates}') + print(f' epsilon: {driver.epsilon}') + +driver.save(driver.save_dir, 'DoubleDQN') +driver.write_log( + episode_date_list, + episode_time_list, + episode_reward_list, + episode_length_list, + episode_loss_list, + episode_epsilon_list, + log_filename='DoubleDQN_log_test.csv' +) +env.close() +plt.ioff() + +# Evaluation Mode +def evaluate_agent(agent, num_episodes=2, render=True): + """Evaluate the trained agent with visualization""" + if render: + env = gym.make("CarRacing-v3", continuous=False, render_mode="human") + else: + env = gym.make("CarRacing-v3", continuous=False, render_mode="rgb_array") + + env = SkipFrame(env, skip=4) # Use SkipFrame from DQN_model.py + env = GrayscaleObservation(env) + env = ResizeObservation(env, (84, 84)) + env = FrameStackObservation(env, stack_size=4) + + agent.epsilon = 0 # Disable exploration + seeds_list = [i for i in range(num_episodes)] + scores = [] + + for episode, seed in enumerate(seeds_list): + state, info = env.reset(seed=seed) + score = 0 + updating = True + + while updating: + action = agent.take_action(state) + state, reward, terminated, truncated, info = env.step(action) + score += reward + updating = not (terminated or truncated) + + scores.append(score) + print(f"Evaluation Episode {episode+1}/{num_episodes} | Seed: {seed} | Score: {score:.1f}") + + env.close() + return np.mean(scores) + +# Run evaluation after training +print("\n=== Starting Evaluation ===") +avg_score = evaluate_agent(driver, num_episodes=2) +print(f"\nAverage evaluation score: {avg_score:.1f}") +plt.show() \ No newline at end of file