| import torch
|
| from torch.utils.data import DataLoader, Subset
|
| from dataset_dt import TrajectoryDataset
|
| from model_dt import DecisionTransformer
|
| import glob
|
| import os
|
| import gc
|
| import time
|
| import argparse
|
| import sys
|
|
|
| def train_sequential_rl_steps_BC(data_prefix="trajectory_data_part_", output_model="dt_model_trained_rl_limited.pth", target_rl_steps=1000000):
|
| log_dir = "./logs"
|
| seq_len = 32
|
| batch_size = 32
|
| learning_rate = 1e-4
|
| epochs = 3
|
|
|
|
|
|
|
|
|
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| print(f"[INFO] Using device: {device}")
|
|
|
|
|
| pattern = f"{data_prefix}*.pkl"
|
| split_files = sorted(glob.glob(pattern))
|
| if not split_files:
|
| split_files = sorted(glob.glob(os.path.join(log_dir, pattern)))
|
|
|
| if not split_files:
|
| print(f"[ERROR] No files found matching prefix: {data_prefix}")
|
| return
|
|
|
| print(f"[INFO] Found {len(split_files)} split files matching '{data_prefix}'.")
|
|
|
|
|
| print("[INFO] Peeking at first file for dimensions...")
|
| temp_dataset = TrajectoryDataset(log_dir, seq_len=seq_len, specific_file=split_files[0])
|
| if len(temp_dataset) > 0:
|
| obs_dim = temp_dataset[0]["observations"].shape[-1]
|
| act_dim = temp_dataset[0]["actions"].shape[-1]
|
| else:
|
|
|
| obs_dim = 9
|
| act_dim = 3
|
|
|
| del temp_dataset
|
| gc.collect()
|
|
|
| print(f"[INFO] Obs Dim: {obs_dim}, Act Dim: {act_dim}")
|
|
|
| model = DecisionTransformer(
|
| obs_dim=obs_dim,
|
| act_dim=act_dim,
|
| hidden=256,
|
| n_layers=4,
|
| n_heads=4,
|
| max_len=4096
|
| ).to(device)
|
|
|
| optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate)
|
|
|
| print(f"[INFO] Starting FRESH training with RL STEP LIMIT: {target_rl_steps} (Pre-loading to Memory)")
|
|
|
|
|
| loaded_datasets = []
|
| total_rl_steps_loaded = 0
|
|
|
| for i, pkl_file in enumerate(split_files):
|
| steps_needed = target_rl_steps - total_rl_steps_loaded
|
| if steps_needed <= 0:
|
| break
|
|
|
| print(f"[INFO] Pre-loading chunk {i+1}/{len(split_files)}: {pkl_file}")
|
| dataset = TrajectoryDataset(log_dir, seq_len=seq_len, specific_file=pkl_file)
|
| dataset_len = len(dataset)
|
|
|
| if dataset_len == 0:
|
| continue
|
|
|
| if dataset_len > steps_needed:
|
| print(f" [LIMIT] Trimming chunk to {steps_needed} samples.")
|
| dataset = Subset(dataset, range(steps_needed))
|
| loaded_datasets.append(dataset)
|
| total_rl_steps_loaded += steps_needed
|
| break
|
| else:
|
| loaded_datasets.append(dataset)
|
| total_rl_steps_loaded += dataset_len
|
|
|
| print(f" [PROGRESS] Memory Buffer: {total_rl_steps_loaded} / {target_rl_steps}")
|
|
|
| if not loaded_datasets:
|
| print("[ERROR] No data loaded! Check file paths.")
|
| return
|
|
|
|
|
| print(f"[INFO] Data Pre-loading Complete. Starting Training on {len(loaded_datasets)} chunks.")
|
|
|
| global_gradient_steps = 0
|
|
|
| for epoch in range(epochs):
|
| print(f"\n=== Epoch {epoch+1}/{epochs} ===")
|
| epoch_start_time = time.time()
|
|
|
| for i, dataset in enumerate(loaded_datasets):
|
|
|
| loader = DataLoader(dataset, batch_size=batch_size, shuffle=True)
|
|
|
| model.train()
|
| chunk_loss = 0.0
|
| chunk_steps = 0
|
|
|
| for batch in loader:
|
| states = batch['observations'].to(device)
|
| actions = batch['actions'].to(device)
|
| returns = batch['returns_to_go'].to(device)
|
| timesteps = batch['timesteps'].to(device)
|
|
|
| returns = torch.zeros_like(returns)
|
|
|
|
|
| action_preds = model(
|
| obs=states,
|
| act=actions,
|
| rtg=returns,
|
| timesteps=timesteps
|
| )
|
|
|
|
|
| action_target = batch['target_actions'].to(device)
|
| loss = torch.mean((action_preds - action_target) ** 2)
|
|
|
| optimizer.zero_grad()
|
| loss.backward()
|
| optimizer.step()
|
|
|
| chunk_loss += loss.item()
|
| chunk_steps += 1
|
| global_gradient_steps += 1
|
|
|
| if chunk_steps % 100 == 0:
|
| print(f" Grad Step {chunk_steps}, Loss: {loss.item():.4f}", end="\r")
|
|
|
| avg_chunk_loss = chunk_loss / chunk_steps if chunk_steps > 0 else 0
|
| print(f" Chunk {i+1} Finished. Avg Loss: {avg_chunk_loss:.4f}")
|
|
|
| print(f"Epoch {epoch+1} completed in {time.time() - epoch_start_time:.2f}s.")
|
|
|
|
|
| dir_name, file_name = os.path.split(output_model)
|
| epoch_model_path = os.path.join(dir_name, f"E_{epoch+1}_{file_name}")
|
| torch.save(model.state_dict(), epoch_model_path)
|
| print(f"[INFO] Saved Epoch {epoch+1} Model to: {epoch_model_path}")
|
|
|
| print(f"\n[DONE] Training Finished.")
|
| print(f" Total RL Steps Processed (cached): {total_rl_steps_loaded}")
|
| print(f" Total Gradient Steps: {global_gradient_steps}")
|
|
|
| print(f"\n[DONE] Training Finished.")
|
| print(f" Total RL Steps Processed: {total_rl_steps_loaded}")
|
| print(f" Total Gradient Steps: {global_gradient_steps}")
|
|
|
|
|
| torch.save(model.state_dict(), output_model)
|
| print(f"[INFO] Saved Final Model to: {output_model}")
|
|
|
| if __name__ == "__main__":
|
| parser = argparse.ArgumentParser(description="Train Decision Transformer with RL Step Limit")
|
| parser.add_argument("--prefix", type=str, default="trajectory_data_part_", help="Prefix of the pickle files to load")
|
| parser.add_argument("--output", type=str, default="dt_model_1M_RL.pth", help="Output filename")
|
| parser.add_argument("--target_rl_steps", type=int, default=1000000, help="Total RL steps (samples) to train on")
|
|
|
| args = parser.parse_args()
|
|
|
| print("\n[Configuration]")
|
| print(f" Prefix: {args.prefix}")
|
| print(f" Output: {args.output}")
|
| print(f" Target RL Steps: {args.target_rl_steps}")
|
| print("-" * 30)
|
|
|
| train_sequential_rl_steps(data_prefix=args.prefix, output_model=args.output, target_rl_steps=args.target_rl_steps)
|
|
|