stargatek1/SSM-MetaRL-Unified
0
1#!/usr/bin/env python32"""3Improved Training Script for SSM-MetaRL-Unified4 5Trains a better-performing model with optimized hyperparameters.6 7Usage:8 python train_improved_model.py --epochs 100 --mode hybrid9"""10 11import argparse12import torch13import numpy as np14from pathlib import Path15import logging16 17from core.ssm import StateSpaceModel18from meta_rl.meta_maml import MetaMAML19from env_runner.environment import Environment20 21# Setup logging22logging.basicConfig(23 level=logging.INFO,24 format='%(asctime)s - %(levelname)s - %(message)s'25)26logger = logging.getLogger(__name__)27 28 29def train_improved_model(30 num_epochs=100,31 tasks_per_epoch=10,32 state_dim=64, # Increased from 3233 hidden_dim=128, # Increased from 6434 inner_lr=0.01,35 outer_lr=0.001,36 adaptation_steps=10, # Increased from 537 mode='hybrid'38):39 """Train an improved model with better hyperparameters"""40 41 logger.info("="*60)42 logger.info("Training Improved SSM-MetaRL Model")43 logger.info("="*60)44 logger.info(f"Epochs: {num_epochs}")45 logger.info(f"Tasks per epoch: {tasks_per_epoch}")46 logger.info(f"State dim: {state_dim}")47 logger.info(f"Hidden dim: {hidden_dim}")48 logger.info(f"Inner LR: {inner_lr}")49 logger.info(f"Outer LR: {outer_lr}")50 logger.info(f"Adaptation steps: {adaptation_steps}")51 logger.info(f"Mode: {mode}")52 logger.info("="*60)53 54 # Create environment55 env = Environment('CartPole-v1')56 57 # Create model with larger capacity58 model = StateSpaceModel(59 state_dim=state_dim,60 input_dim=4, # CartPole observation space61 output_dim=4, # Output dimension for SSM62 hidden_dim=hidden_dim63 )64 65 logger.info(f"Model created with {sum(p.numel() for p in model.parameters())} parameters")66 67 # Create meta-learner68 meta_learner = MetaMAML(69 model=model,70 inner_lr=inner_lr,71 outer_lr=outer_lr,72 adaptation_steps=adaptation_steps73 )74 75 logger.info("Starting meta-training...")76 77 # Training loop78 best_reward = 079 for epoch in range(num_epochs):80 epoch_losses = []81 epoch_rewards = []82 83 for task_idx in range(tasks_per_epoch):84 # Generate task data85 support_obs, support_actions, query_obs, query_actions = generate_task_data(86 env, model, num_episodes=387 )88 89 # Meta-training step90 loss = meta_learner.meta_train_step(91 support_x=support_obs,92 support_y=support_actions,93 query_x=query_obs,94 query_y=query_actions95 )96 97 epoch_losses.append(loss)98 99 # Evaluate on a test episode100 reward = evaluate_model(env, model)101 epoch_rewards.append(reward)102 103 avg_loss = np.mean(epoch_losses)104 avg_reward = np.mean(epoch_rewards)105 106 if (epoch + 1) % 10 == 0 or epoch == 0:107 logger.info(f"Epoch {epoch+1}/{num_epochs}: Loss = {avg_loss:.4f}, Reward = {avg_reward:.2f}")108 109 # Save best model110 if avg_reward > best_reward:111 best_reward = avg_reward112 model.save(f'models/cartpole_{mode}_improved_model.pth')113 logger.info(f" → New best model saved! Reward: {avg_reward:.2f}")114 115 logger.info("="*60)116 logger.info(f"Training complete! Best reward: {best_reward:.2f}")117 logger.info(f"Model saved to: models/cartpole_{mode}_improved_model.pth")118 logger.info("="*60)119 120 return model121 122 123def generate_task_data(env, model, num_episodes=3):124 """Generate training data from environment episodes"""125 all_obs = []126 all_actions = []127 128 for _ in range(num_episodes):129 obs = env.reset()130 hidden = model.init_hidden(batch_size=1)131 132 episode_obs = []133 episode_actions = []134 135 for step in range(200): # Max 200 steps per episode136 obs_tensor = torch.FloatTensor(obs).unsqueeze(0)137 138 with torch.no_grad():139 action_logits, hidden = model(obs_tensor, hidden)140 141 # Use first 2 dimensions for action (CartPole has 2 actions)142 action = torch.argmax(action_logits[:, :2], dim=-1).item()143 144 episode_obs.append(obs)145 episode_actions.append(action)146 147 next_obs, reward, done, info = env.step(action)148 149 if done:150 break151 152 obs = next_obs153 154 all_obs.extend(episode_obs)155 all_actions.extend(episode_actions)156 157 # Convert to tensors158 obs_tensor = torch.FloatTensor(all_obs).unsqueeze(0) # (1, T, 4)159 actions_tensor = torch.LongTensor(all_actions).unsqueeze(0).unsqueeze(-1).float() # (1, T, 1)160 161 # Split into support and query sets162 split_idx = len(all_obs) // 2163 support_obs = obs_tensor[:, :split_idx, :]164 support_actions = actions_tensor[:, :split_idx, :]165 query_obs = obs_tensor[:, split_idx:, :]166 query_actions = actions_tensor[:, split_idx:, :]167 168 return support_obs, support_actions, query_obs, query_actions169 170 171def evaluate_model(env, model, num_episodes=5):172 """Evaluate model performance"""173 total_rewards = []174 175 for _ in range(num_episodes):176 obs = env.reset()177 hidden = model.init_hidden(batch_size=1)178 episode_reward = 0179 180 for step in range(500):181 obs_tensor = torch.FloatTensor(obs).unsqueeze(0)182 183 with torch.no_grad():184 action_logits, hidden = model(obs_tensor, hidden)185 186 # Use first 2 dimensions for action187 action = torch.argmax(action_logits[:, :2], dim=-1).item()188 189 next_obs, reward, done, info = env.step(action)190 episode_reward += reward191 192 if done:193 break194 195 obs = next_obs196 197 total_rewards.append(episode_reward)198 199 return np.mean(total_rewards)200 201 202if __name__ == '__main__':203 parser = argparse.ArgumentParser(description='Train improved SSM-MetaRL model')204 parser.add_argument('--epochs', type=int, default=100, help='Number of training epochs')205 parser.add_argument('--tasks_per_epoch', type=int, default=10, help='Tasks per epoch')206 parser.add_argument('--state_dim', type=int, default=64, help='State dimension')207 parser.add_argument('--hidden_dim', type=int, default=128, help='Hidden dimension')208 parser.add_argument('--inner_lr', type=float, default=0.01, help='Inner loop learning rate')209 parser.add_argument('--outer_lr', type=float, default=0.001, help='Outer loop learning rate')210 parser.add_argument('--adaptation_steps', type=int, default=10, help='Adaptation steps')211 parser.add_argument('--mode', type=str, default='hybrid', choices=['standard', 'hybrid'],212 help='Adaptation mode')213 214 args = parser.parse_args()215 216 # Create models directory217 Path('models').mkdir(exist_ok=True)218 Path('logs').mkdir(exist_ok=True)219 220 # Train model221 model = train_improved_model(222 num_epochs=args.epochs,223 tasks_per_epoch=args.tasks_per_epoch,224 state_dim=args.state_dim,225 hidden_dim=args.hidden_dim,226 inner_lr=args.inner_lr,227 outer_lr=args.outer_lr,228 adaptation_steps=args.adaptation_steps,229 mode=args.mode230 )231 232 # Final evaluation233 env = Environment('CartPole-v1')234 final_reward = evaluate_model(env, model, num_episodes=20)235 logger.info(f"\nFinal evaluation (20 episodes): {final_reward:.2f} ± {np.std([evaluate_model(env, model, 1) for _ in range(20)]):.2f}")236 237 # Save training log238 with open(f'logs/training_improved_{args.mode}.log', 'w') as f:239 f.write(f"Training completed\n")240 f.write(f"Final reward: {final_reward:.2f}\n")241 f.write(f"Model: models/cartpole_{args.mode}_improved_model.pth\n")242 243 