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
        }