File size: 7,122 Bytes
b685739
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import torch
from torch.utils.data import DataLoader, Subset
from dataset_dt import TrajectoryDataset
from model_dt import DecisionTransformer
import glob
import os
import gc
import time
import argparse
import sys

def train_sequential_rl_steps_BC(data_prefix="trajectory_data_part_", output_model="dt_model_trained_rl_limited.pth", target_rl_steps=1000000):
    log_dir = "./logs"
    seq_len = 32
    batch_size = 32
    learning_rate = 1e-4
    epochs = 3 
    # Note: Epochs are still iterated, but we will likely break early in the first epoch 
    # if target_rl_steps is reached. If target is huge, it might run multiple epochs.
    # But for "1M RL Steps" on a dataset of >2M, it will stop in Epoch 1.
    
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    print(f"[INFO] Using device: {device}")

    # 1. Initialize Model Dimensions
    pattern = f"{data_prefix}*.pkl"
    split_files = sorted(glob.glob(pattern))
    if not split_files:
        split_files = sorted(glob.glob(os.path.join(log_dir, pattern)))
        
    if not split_files:
        print(f"[ERROR] No files found matching prefix: {data_prefix}")
        return

    print(f"[INFO] Found {len(split_files)} split files matching '{data_prefix}'.")
    
    # Peek at first file for dimensions
    print("[INFO] Peeking at first file for dimensions...")
    temp_dataset = TrajectoryDataset(log_dir, seq_len=seq_len, specific_file=split_files[0])
    if len(temp_dataset) > 0:
        obs_dim = temp_dataset[0]["observations"].shape[-1]
        act_dim = temp_dataset[0]["actions"].shape[-1]
    else:
        # Fallback defaults
        obs_dim = 9
        act_dim = 3
        
    del temp_dataset
    gc.collect()
    
    print(f"[INFO] Obs Dim: {obs_dim}, Act Dim: {act_dim}")

    model = DecisionTransformer(
        obs_dim=obs_dim,
        act_dim=act_dim,
        hidden=256,
        n_layers=4,
        n_heads=4,
        max_len=4096 
    ).to(device)

    optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate)
    
    print(f"[INFO] Starting FRESH training with RL STEP LIMIT: {target_rl_steps} (Pre-loading to Memory)")

    # 1. Pre-load Data Phase
    loaded_datasets = []
    total_rl_steps_loaded = 0
    
    for i, pkl_file in enumerate(split_files):
        steps_needed = target_rl_steps - total_rl_steps_loaded
        if steps_needed <= 0:
            break
            
        print(f"[INFO] Pre-loading chunk {i+1}/{len(split_files)}: {pkl_file}")
        dataset = TrajectoryDataset(log_dir, seq_len=seq_len, specific_file=pkl_file)
        dataset_len = len(dataset)
        
        if dataset_len == 0:
            continue
            
        if dataset_len > steps_needed:
            print(f"  [LIMIT] Trimming chunk to {steps_needed} samples.")
            dataset = Subset(dataset, range(steps_needed))
            loaded_datasets.append(dataset)
            total_rl_steps_loaded += steps_needed
            break
        else:
            loaded_datasets.append(dataset)
            total_rl_steps_loaded += dataset_len
            
        print(f"  [PROGRESS] Memory Buffer: {total_rl_steps_loaded} / {target_rl_steps}")

    if not loaded_datasets:
        print("[ERROR] No data loaded! Check file paths.")
        return

    # 2. Training Phase (Sequential per Chunk, just like original)
    print(f"[INFO] Data Pre-loading Complete. Starting Training on {len(loaded_datasets)} chunks.")
    
    global_gradient_steps = 0
    
    for epoch in range(epochs):
        print(f"\n=== Epoch {epoch+1}/{epochs} ===")
        epoch_start_time = time.time()
        
        for i, dataset in enumerate(loaded_datasets):
            # Same Batch/Shuffle logic as original
            loader = DataLoader(dataset, batch_size=batch_size, shuffle=True)
            
            model.train()
            chunk_loss = 0.0
            chunk_steps = 0
            
            for batch in loader:
                states = batch['observations'].to(device)
                actions = batch['actions'].to(device)
                returns = batch['returns_to_go'].to(device)
                timesteps = batch['timesteps'].to(device)

                returns = torch.zeros_like(returns)
                
                # Forward
                action_preds = model(
                    obs=states,
                    act=actions,
                    rtg=returns,
                    timesteps=timesteps
                )
                
                # Loss
                action_target = batch['target_actions'].to(device)
                loss = torch.mean((action_preds - action_target) ** 2)
                
                optimizer.zero_grad()
                loss.backward()
                optimizer.step()
                
                chunk_loss += loss.item()
                chunk_steps += 1
                global_gradient_steps += 1
                
                if chunk_steps % 100 == 0:
                    print(f"  Grad Step {chunk_steps}, Loss: {loss.item():.4f}", end="\r")
            
            avg_chunk_loss = chunk_loss / chunk_steps if chunk_steps > 0 else 0
            print(f"  Chunk {i+1} Finished. Avg Loss: {avg_chunk_loss:.4f}")

        print(f"Epoch {epoch+1} completed in {time.time() - epoch_start_time:.2f}s.")
        
        # Save Model per Epoch
        dir_name, file_name = os.path.split(output_model)
        epoch_model_path = os.path.join(dir_name, f"E_{epoch+1}_{file_name}")
        torch.save(model.state_dict(), epoch_model_path)
        print(f"[INFO] Saved Epoch {epoch+1} Model to: {epoch_model_path}")
        
    print(f"\n[DONE] Training Finished.")
    print(f"  Total RL Steps Processed (cached): {total_rl_steps_loaded}")
    print(f"  Total Gradient Steps: {global_gradient_steps}")
        
    print(f"\n[DONE] Training Finished.")
    print(f"  Total RL Steps Processed: {total_rl_steps_loaded}")
    print(f"  Total Gradient Steps: {global_gradient_steps}")
        
    # Save Final Model
    torch.save(model.state_dict(), output_model)
    print(f"[INFO] Saved Final Model to: {output_model}")

if __name__ == "__main__":
    parser = argparse.ArgumentParser(description="Train Decision Transformer with RL Step Limit")
    parser.add_argument("--prefix", type=str, default="trajectory_data_part_", help="Prefix of the pickle files to load")
    parser.add_argument("--output", type=str, default="dt_model_1M_RL.pth", help="Output filename")
    parser.add_argument("--target_rl_steps", type=int, default=1000000, help="Total RL steps (samples) to train on")
    
    args = parser.parse_args()
    
    print("\n[Configuration]")
    print(f"  Prefix: {args.prefix}")
    print(f"  Output: {args.output}")
    print(f"  Target RL Steps: {args.target_rl_steps}")
    print("-" * 30)
    
    train_sequential_rl_steps(data_prefix=args.prefix, output_model=args.output, target_rl_steps=args.target_rl_steps)