code3939's picture
Upload 765 files
b7e9b58 verified
Raw History Blame Contribute Delete
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
)