Download Upload/01_Source_Code/Python_Training/finetuning_ext_RLStep.py from code3939/DecisionTransformer-Unity-Sim: direct link, hf CLI and curl.
- Browser
- Download file 8.24 kB
-
https://huggingface.co/code3939/DecisionTransformer-Unity-Sim/resolve/main/Upload/01_Source_Code/Python_Training/finetuning_ext_RLStep.py
- Command line
-
hf download hf://code3939/DecisionTransformer-Unity-Sim/Upload/01_Source_Code/Python_Training/finetuning_ext_RLStep.py
-
curl -L -o finetuning_ext_RLStep.py https://huggingface.co/code3939/DecisionTransformer-Unity-Sim/resolve/main/Upload/01_Source_Code/Python_Training/finetuning_ext_RLStep.py
8.24 kB
| 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 | |
| ) | |