# 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 }