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 # Note: Epochs are still iterated, but we will likely break early in the first epoch # if target_rl_steps is reached. If target is huge, it might run multiple epochs. # But for "1M RL Steps" on a dataset of >2M, it will stop in Epoch 1. device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"[INFO] Using device: {device}") # 1. Initialize Model Dimensions 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}'.") # Peek at first file for dimensions 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: # Fallback defaults 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)") # 1. Pre-load Data Phase 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 # 2. Training Phase (Sequential per Chunk, just like original) 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): # Same Batch/Shuffle logic as original 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) # Forward action_preds = model( obs=states, act=actions, rtg=returns, timesteps=timesteps ) # Loss 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.") # Save Model per Epoch 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}") # Save Final Model 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)