File size: 7,122 Bytes
b685739 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 | 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)
|