File size: 8,046 Bytes
b7e9b58 | 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 185 186 187 188 189 190 191 192 193 194 | # dataset_dt.py
import json
import torch
from torch.utils.data import Dataset
import glob
import os
import pickle
import numpy as np
class TrajectoryDataset(Dataset):
def __init__(self, log_dir, seq_len=32, specific_file=None):
self.seq_len = seq_len
self.raw_data = []
# NEW APPROACH: Load data into EPISODES.
raw_rows = []
if specific_file:
print(f"[INFO] Loading specific file: {specific_file}...")
try:
with open(specific_file, "rb") as f:
raw_rows = pickle.load(f)
print(f"[INFO] Loaded {len(raw_rows)} steps from {specific_file}.")
except (EOFError, pickle.UnpicklingError) as e:
print(f"[ERROR] Failed to load {specific_file}: {e}")
raw_rows = []
else:
# Check for split pickle files first (Memory Efficient Load)
split_files = sorted(glob.glob(os.path.join(log_dir, "trajectory_data_part_*.pkl")))
if not split_files:
split_files = sorted(glob.glob("trajectory_data_part_*.pkl"))
if split_files:
raise RuntimeError(
f"[ERROR] Found {len(split_files)} split files but 'specific_file' was not provided.\n"
"Cannot load all split files at once due to memory constraints.\n"
"Please use 'train_sequential.py' to train sequentially on chunks."
)
else:
# Fallback to single pickle or JSON
pickle_file = "trajectory_data.pkl"
if os.path.exists(pickle_file):
print(f"[INFO] Loading data from {pickle_file}...")
try:
with open(pickle_file, "rb") as f:
raw_rows = pickle.load(f)
print(f"[INFO] Loaded {len(raw_rows)} steps from pickle.")
except Exception as e:
print(f"[WARNING] Pickle corrupted: {e}")
raw_rows = []
if not raw_rows:
json_files = sorted(glob.glob(os.path.join(log_dir, "*.json")))
if json_files:
print(f"[INFO] Loading {len(json_files)} JSON files...")
for jf in json_files:
try:
with open(jf, 'r') as f:
data = json.load(f)
rows = data.get("data", [])
# Normalize JSON rows to lists if needed
for r in rows:
vals = r["values"] if isinstance(r, dict) else r
if isinstance(vals, list) and len(vals) == 14:
raw_rows.append(vals)
except: pass
# Process raw_rows into Episodes => Steps
# We need to reconstruct episodes to calculate RTG correctly.
self.episodes = []
current_episode = []
for row in raw_rows:
# row can be a dict (from convert_added_json_to_pickle) OR a list (legacy raw values)
if isinstance(row, dict):
# Format: {'obs': tensor, 'act': tensor, 'rew': tensor, 'done': tensor}
obs = row['obs']
act = row['act']
rew = row['rew']
done = row['done']
# Ensure they are tensors
if not isinstance(obs, torch.Tensor): obs = torch.tensor(obs, dtype=torch.float32)
if not isinstance(act, torch.Tensor): act = torch.tensor(act, dtype=torch.float32)
if not isinstance(rew, torch.Tensor): rew = torch.tensor(rew, dtype=torch.float32)
if not isinstance(done, torch.Tensor): done = torch.tensor(done, dtype=torch.float32)
else:
# Legacy List format
vals = torch.tensor(row, dtype=torch.float32)
obs = vals[:9]
act_cont = vals[9:11]
act_fire = vals[11:12]
rew = vals[12:13]
done = vals[13:14]
act = torch.cat([act_cont, act_fire])
step = {
"obs": obs,
"act": act,
"rew": rew,
"done": done
}
current_episode.append(step)
# Check done flag (assuming scalar tensor)
if done.item() > 0.5:
self.episodes.append(self._process_episode(current_episode))
current_episode = []
# Handle trailing data
if current_episode:
self.episodes.append(self._process_episode(current_episode))
# Flatten for indexing
self.indices = []
for ep_idx, ep in enumerate(self.episodes):
length = len(ep["obs"])
for t in range(length):
self.indices.append((ep_idx, t))
print(f"[INFO] Processed {len(self.episodes)} episodes. Total {len(self.indices)} samples.")
def _process_episode(self, episode_steps):
# Calculate RTG for this episode
rews = torch.stack([s["rew"] for s in episode_steps]) # (L, 1)
# RTG calculation: Cumulative sum from back to front
rtg = torch.flip(torch.cumsum(torch.flip(rews, dims=[0]), dim=0), dims=[0])
if rtg.dim() == 1:
rtg = rtg.unsqueeze(-1) # Ensure (L, 1)
# Timesteps
timesteps = torch.arange(len(episode_steps), dtype=torch.long)
# Stack others
obs = torch.stack([s["obs"] for s in episode_steps])
acts = torch.stack([s["act"] for s in episode_steps])
dones = torch.stack([s["done"] for s in episode_steps])
# Create Shifted Actions (a_{t-1}) for input
# Prepend zero, remove last
first_zero = torch.zeros(1, acts.shape[1])
shifted_acts = torch.cat([first_zero, acts[:-1]], dim=0)
return {
"obs": obs,
"act": acts, # Target (a_t)
"shifted_act": shifted_acts, # Input (a_{t-1})
"rtg": rtg,
"timesteps": timesteps,
"len": len(episode_steps)
}
def __len__(self):
return len(self.indices)
def __getitem__(self, idx):
ep_idx, start_t = self.indices[idx]
ep = self.episodes[ep_idx]
end_t = start_t + self.seq_len
real_end = min(end_t, ep["len"])
# Aligned Slicing: All tensors sliced identically [start : end]
# Logic: Input(s_t, a_{t-1}) -> Target(a_t)
obs = ep["obs"][start_t : real_end]
rtg = ep["rtg"][start_t : real_end]
time = ep["timesteps"][start_t : real_end]
target = ep["act"][start_t : real_end] # a_t (Target)
act = ep["shifted_act"][start_t : real_end] # a_{t-1} (Input)
# Pad if necessary to match seq_len
# Pad to full seq_len
if obs.shape[0] < self.seq_len:
pad = self.seq_len - obs.shape[0]
obs = torch.cat([obs, torch.zeros(pad, obs.shape[1])])
act = torch.cat([act, torch.zeros(pad, act.shape[1])])
target = torch.cat([target, torch.zeros(pad, act.shape[1])])
rtg = torch.cat([rtg, torch.zeros(pad, 1)])
time = torch.cat([time, torch.zeros(pad, dtype=torch.long)])
return {
"observations": obs,
"actions": act, # Input a_{t-1}
"returns_to_go": rtg,
"timesteps": time,
"target_actions": target # Target a_t
}
|