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 import re def finetuning_rl_steps( data_prefix="trajectory_data_part_", output_model="dt_model_finetuned.pth", load_model="dt_model_trained.pth", target_rl_steps=1000000, epochs=5, learning_rate=1e-5 ): log_dir = "./logs" seq_len = 32 batch_size = 32 device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"[INFO] Using device: {device}") # 1. Initialize Model Dimensions (Peek Logic) 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: # Fallback for spelling if "Trajectoy" in data_prefix: fallback = data_prefix.replace("Trajectoy", "Trajectory") split_files = sorted(glob.glob(fallback)) 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}") # 2. Model Setup model = DecisionTransformer( obs_dim=obs_dim, act_dim=act_dim, hidden=256, n_layers=4, n_heads=4, max_len=4096 ).to(device) # Load Checkpoint if os.path.exists(load_model): print(f"[INFO] Loading checkpoint: {load_model}") try: state_dict = torch.load(load_model, map_location=device) model.load_state_dict(state_dict) print("[INFO] Model loaded successfully.") except Exception as e: print(f"[ERROR] Failed to load checkpoint: {e}") return else: print(f"[WARNING] Checkpoint '{load_model}' not found! Starting FRESH training (Fine-Tuning aborted).") # Depending on user intent, we might want to return here. # But usually we proceed if user insists, though it's technically Pre-training then. print(f"[INFO] Learning rate: {learning_rate}") optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate) print(f"[INFO] Starting FINE-TUNING with RL STEP LIMIT: {target_rl_steps} (Pre-loading to Memory)") # 3. 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 # 4. Training Phase (Sequential Chunk Processing) print(f"[INFO] Data Pre-loading Complete. Starting Fine-Tuning on {len(loaded_datasets)} chunks.") global_gradient_steps = 0 for epoch in range(epochs): print(f"\n=== Fine-Tuning 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) # 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} Checkpoint to: {epoch_model_path}") print(f"\n[DONE] Fine-Tuning Finished.") print(f" Total RL Steps Processed (cached): {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="Fine-Tune 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_finetuned.pth", help="Output filename") parser.add_argument("--load_model", type=str, required=True, help="Path to pre-trained model checkpoint") parser.add_argument("--target_rl_steps", type=int, default=1000000, help="Total RL steps (samples) to train on") parser.add_argument("--epochs", type=int, default=5, help="Number of epochs") parser.add_argument("--learning_rate", type=float, default=1e-5, help="Learning rate (default: 1e-5)") args = parser.parse_args() target_steps = args.target_steps if args.target_steps > 0 else 1000000 print("\n[Configuration]") print(f" Load Model: {args.load_model}") print(f" Prefix: {args.prefix}") print(f" Output: {args.output}") print(f" Target RL Steps: {target_steps}") print(f" Epochs: {args.epochs}") print(f" LR: {args.learning_rate}") print("-" * 30) finetuning_rl_steps( data_prefix=args.prefix, output_model=args.output, load_model=args.load_model, target_rl_steps=target_steps, epochs=args.epochs, learning_rate=args.learning_rate )