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