DecisionTransformer-Unity-Sim / DT Update /MakingBC /train_sequential_ext_RLStep_For_BC.py
code3939's picture
Upload 28 files
b685739 verified
Raw
History Blame Contribute Delete
7.12 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
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)