Download Upload/01_Source_Code/Python_Training/dataset_dt.py from code3939/DecisionTransformer-Unity-Sim: direct link, hf CLI and curl.
- Browser
- Download file 8.05 kB
-
https://huggingface.co/code3939/DecisionTransformer-Unity-Sim/resolve/main/Upload/01_Source_Code/Python_Training/dataset_dt.py
- Command line
-
hf download hf://code3939/DecisionTransformer-Unity-Sim/Upload/01_Source_Code/Python_Training/dataset_dt.py
-
curl -L -o dataset_dt.py https://huggingface.co/code3939/DecisionTransformer-Unity-Sim/resolve/main/Upload/01_Source_Code/Python_Training/dataset_dt.py
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 | |
| } | |