CoolFace
Apppublic

stargatek1/SSM-MetaRL-Unified

sourceHugging Facemitupdated 11mo agoView on Hugging Face
0likes
train_improved_model.py243 linesDownload Raw Back to root
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