Add more hyperparameter scan modes

This commit is contained in:
2025-11-25 09:38:25 +01:00
parent b012a86768
commit c0b2545671
5 changed files with 1152 additions and 112 deletions
File diff suppressed because one or more lines are too long
+268 -73
View File
@@ -1,90 +1,285 @@
import torch
from aiRNN import dataloader, models, losses
import pathlib
import pandas as pd
from aiRNN import dataloader, models, losses
file_list = pathlib.Path("../../cpp/known_routes_and_aircraft.csv")
base_path = file_list.parent
file_list = file_list.read_text().splitlines()
file_list = [(base_path / f).resolve() for f in file_list if (base_path / f).exists()]
MODE = 1
DEVICE = "cuda"
if not pathlib.Path("dataset.pt").exists():
dataset = dataloader.EvenlySpacedDataset(
filepaths=file_list,
n_input=30*15, # 15 minutes input
n_output=30*5, # 5 minutes output
n_windows_per_file=7,
step=1,
feature_columns=("lat", "lon", "alt", "ias"),
context_columns=("last_lat", "last_lon", "last_alt", "last_ias", "last_timestamp"),
time_columns=("timestamp", "dt"),
target_columns=("lat", "lon", "alt"),
device="cuda",
)
dataset.save_entire_dataset("dataset.pt")
results = []
for step in [1, 5, 10, 30]:
dataset = dataloader.SaveDataset(torch.load("dataset.pt"), device="cuda", step=step)
dataset_length = len(dataset)
dataset, val_dataset = torch.utils.data.random_split(
# ------------------------------------------------------------
# Dataset creation
# ------------------------------------------------------------
def load_file_list(path: pathlib.Path):
base = path.parent
return [
(base / f).resolve()
for f in path.read_text().splitlines()
if (base / f).exists()
]
def ensure_dataset(path, **kwargs):
if not pathlib.Path(path).exists():
ds = dataloader.EvenlySpacedDataset(**kwargs, device=DEVICE)
ds.save_entire_dataset(path)
file_list_path = pathlib.Path("../../cpp/known_routes_and_aircraft.csv")
file_list = load_file_list(file_list_path)
ensure_dataset(
"dataset.pt",
filepaths=file_list,
n_input=30 * 15,
n_output=30 * 5,
n_windows_per_file=7,
step=1,
feature_columns=("lat", "lon", "alt", "ias"),
context_columns=("last_lat", "last_lon", "last_alt", "last_ias", "last_timestamp"),
time_columns=("timestamp", "dt"),
target_columns=("lat", "lon", "alt"),
)
ensure_dataset(
"long_dataset.pt",
filepaths=file_list,
n_input=30 * 60,
n_output=30 * 60,
n_windows_per_file=5,
step=1,
feature_columns=("lat", "lon", "alt", "ias"),
context_columns=("last_lat", "last_lon", "last_alt", "last_ias", "last_timestamp"),
time_columns=("timestamp", "dt"),
target_columns=("lat", "lon", "alt"),
)
# ------------------------------------------------------------
# Training utilities
# ------------------------------------------------------------
def split_dataset(dataset, frac=0.8, seed=42):
n = len(dataset)
train_n = int(frac * n)
return torch.utils.data.random_split(
dataset,
[int(0.8 * dataset_length), dataset_length - int(0.8 * dataset_length)],
generator=torch.Generator().manual_seed(42)
[train_n, n - train_n],
generator=torch.Generator().manual_seed(seed),
)
print(f"Starting hyperparameter scan for step={step}")
for base_name, base_model in [
("RNN", models.ThreeInputRNN),
("LSTM", models.ThreeInputLSTM),
("GRU", models.ThreeInputGRU),
]:
for hidden_size in [16, 32, 64]:
for rnn_size in [32, 64, 128]:
model = base_model(
def train_one_model(model, dataset, val_dataset, warm_up_steps, pred_steps, epochs=100, altitude_weight=1e-3):
optimizer = torch.optim.Adam(model.parameters(), lr=1e-2)
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, patience=5, factor=0.5)
criterion = losses.HaversineAltitudeLoss(alt_const=altitude_weight)
for epoch in range(epochs):
batch_losses = []
loader = torch.utils.data.DataLoader(dataset, batch_size=256, shuffle=True, collate_fn=dataloader.collate_to_cuda(DEVICE), num_workers=4)
for X_f, X_t, y, X_c in loader:
optimizer.zero_grad()
y_pred, _ = model(X_t, X_f, X_c, warm_up_steps, pred_steps)
loss = criterion(y_pred, y)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()
batch_losses.append(loss.item())
mean_loss = sum(batch_losses) / len(batch_losses)
scheduler.step(mean_loss)
print(f"Epoch {epoch}: loss={mean_loss}, lr={optimizer.param_groups[0]['lr']}")
# Validation
with torch.no_grad():
val_losses = []
val_loader = torch.utils.data.DataLoader(val_dataset, batch_size=256, collate_fn=dataloader.collate_to_cuda(DEVICE), num_workers=4)
for X_f, X_t, y, X_c in val_loader:
y_pred, _ = model(X_t, X_f, X_c, warm_up_steps, pred_steps)
val_losses.append(criterion(y_pred, y).item())
val_loss = sum(val_losses) / len(val_losses)
print(f"Validation loss: {val_loss}")
return val_loss
def save_results(results, preliminary=False):
if MODE == 1:
cols = ["Model", "Hidden Size", "RNN Size", "Step", "Final Loss"]
elif MODE == 2:
cols = ["Model", "Hidden Layers", "RNN Layers", "Dropout", "Final Loss"]
elif MODE == 3:
cols = ["Model", "Warm-up Steps", "Prediction Steps", "Altitude Weight", "Final Loss"]
else:
raise ValueError("Invalid MODE")
df = pd.DataFrame(results, columns=cols)
suffix = "_preliminary" if preliminary else ""
df.to_csv(f"hyperparameter_scan_results_mode{MODE}{suffix}.csv", index=False)
# ------------------------------------------------------------
# MODE 1 sweep
# ------------------------------------------------------------
def run_mode_1():
results = []
for step in [1, 5, 10, 30]:
base_ds = dataloader.SaveDataset(torch.load("dataset.pt"), step=step)
train_ds, val_ds = split_dataset(base_ds)
print(f"Starting hyperparameter scan for step={step}")
for name, cls in [("RNN", models.ThreeInputRNN),
("LSTM", models.ThreeInputLSTM),
("GRU", models.ThreeInputGRU)]:
for hidden_size in [16, 32, 64]:
for rnn_size in [32, 64, 128]:
model = cls(
time_in=2,
feat_in=4,
context_in=5,
hidden_size=hidden_size,
rnn_size=rnn_size,
out_size=3,
device=DEVICE,
)
warm = 450 // step
pred = 150 // step
print(f"Training {name} hs={hidden_size} rs={rnn_size}")
loss = train_one_model(model, train_ds, val_ds, warm, pred)
torch.save(model.state_dict(), f"{name}_hs{hidden_size}_rs{rnn_size}_step{step}.pt")
results.append((name, hidden_size, rnn_size, step, loss))
save_results(results, preliminary=True)
return results
# ------------------------------------------------------------
# MODE 2 sweep
# ------------------------------------------------------------
def run_mode_2():
results = []
step = 1
base_ds = dataloader.SaveDataset(torch.load("dataset.pt"), step=step)
train_ds, val_ds = split_dataset(base_ds)
base_name = "GRU"
cls = models.ThreeInputGRU
hidden_size = 32
rnn_size = 64
for hl in [0, 1, 2, 4]:
for rl in [2, 3, 4, 5]:
for dropout in [0.0, 0.1, 0.2, 0.3]:
model = cls(
time_in=2,
feat_in=4,
context_in=5,
hidden_size=hidden_size,
rnn_size=rnn_size,
out_size=3,
device="cuda",
hidden_layers=hl,
rnn_layers=rl,
rnn_dropout=dropout,
device=DEVICE,
)
print(f"Training {base_name} with hidden_size={hidden_size}, rnn_size={rnn_size}")
optimizer = torch.optim.Adam(model.parameters(), lr=1e-2)
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, patience=5, factor=0.5)
criterion = losses.HaversineAltitudeLoss(alt_const=1e-3)
loss = -1.0
for epoch in range(100):
loss_history = []
for X_f, X_t, y, X_c in torch.utils.data.DataLoader(dataset, batch_size=256, shuffle=True):
optimizer.zero_grad()
warm_up_steps = 450 // step
pred_steps = 150 // step
y_pred, _ = model(X_t, X_f, X_c, warm_up_steps, pred_steps)
loss = criterion(y_pred, y)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()
loss_history.append(loss.item())
loss = sum(loss_history) / len(loss_history)
scheduler.step(loss)
print(f"Epoch {epoch}: loss={loss}, lr={optimizer.param_groups[0]['lr']}")
with torch.no_grad():
val_loss_history = []
for X_f, X_t, y, X_c in torch.utils.data.DataLoader(val_dataset, batch_size=256):
warm_up_steps = 450 // step
pred_steps = 150 // step
y_pred, _ = model(X_t, X_f, X_c, warm_up_steps, pred_steps)
v_loss = criterion(y_pred, y)
val_loss_history.append(v_loss.item())
loss = sum(val_loss_history) / len(val_loss_history)
print(f"Validation loss: {loss}")
torch.save(model.state_dict(), f"{base_name}_hs{hidden_size}_rs{rnn_size}_step{step}.pt")
results.append((base_name, hidden_size, rnn_size, step, loss))
for r in results:
print(f"Model: {r[0]}, hidden_size={r[1]}, rnn_size={r[2]}, step={r[3]} => final loss={r[4]}")
warm = 450 // step
pred = 150 // step
# Save results to CSV
df = pd.DataFrame(results, columns=["Model", "Hidden Size", "RNN Size", "Step", "Final Loss"])
df.to_csv("hyperparameter_scan_results.csv", index=False)
print(f"Training {base_name} hl={hl} rl={rl} do={dropout}")
loss = train_one_model(model, train_ds, val_ds, warm, pred)
torch.save(model.state_dict(), f"{base_name}_hl{hl}_rl{rl}_do{int(dropout*10)}.pt")
results.append((base_name, hl, rl, dropout, loss))
save_results(results, preliminary=True)
return results
# ------------------------------------------------------------
# MODE 3 sweep
# ------------------------------------------------------------
def run_mode_3():
results = []
step = 1
base_name = "GRU"
cls = models.ThreeInputGRU
hidden_size = 32
rnn_size = 64
hidden_layers = 1
rnn_layers = 3
rnn_dropout = 0.1
for warm in [300, 600, 900, 1800]:
for pred in [150, 300, 600, 900, 1800]:
start = 30 * 60 - warm
end = 30 * 60 - pred
base_ds = dataloader.SaveDataset(
torch.load("long_dataset.pt"),
step=step,
start_offset=start,
end_offset=end,
)
train_ds, val_ds = split_dataset(base_ds)
for altitude_weight in [1e-5, 1e-4, 1e-3]:
model = cls(
time_in=2,
feat_in=4,
context_in=5,
hidden_size=hidden_size,
rnn_size=rnn_size,
out_size=3,
hidden_layers=hidden_layers,
rnn_layers=rnn_layers,
rnn_dropout=rnn_dropout,
device=DEVICE,
)
print(f"Training {base_name} warm={warm} pred={pred} altitude_weight={altitude_weight}")
loss = train_one_model(model, train_ds, val_ds, warm // step, pred // step, altitude_weight=altitude_weight)
torch.save(model.state_dict(), f"{base_name}_wu{warm}_ps{pred}_aw{altitude_weight}.pt")
results.append((base_name, warm, pred, altitude_weight, loss))
save_results(results, preliminary=True)
return results
# ------------------------------------------------------------
# Dispatch
# ------------------------------------------------------------
if MODE == 1:
results = run_mode_1()
elif MODE == 2:
results = run_mode_2()
elif MODE == 3:
results = run_mode_3()
else:
raise ValueError("Invalid MODE.")
save_results(results)
@@ -0,0 +1,109 @@
Model,Hidden Size,RNN Size,Step,Final Loss
RNN,16,32,1,372.1873533829399
RNN,16,64,1,371.67606353759766
RNN,16,128,1,328.19922903309697
RNN,32,32,1,376.71770842179006
RNN,32,64,1,369.6340640524159
RNN,32,128,1,281.8589802617612
RNN,64,32,1,378.75125188412875
RNN,64,64,1,405.08673095703125
RNN,64,128,1,1088.4633191979449
LSTM,16,32,1,152.01849580847698
LSTM,16,64,1,152.46532937754756
LSTM,16,128,1,39.81503051260243
LSTM,32,32,1,45.500295265861176
LSTM,32,64,1,51.218478575996734
LSTM,32,128,1,32.29804623645285
LSTM,64,32,1,105.44863045733908
LSTM,64,64,1,62.87264318051545
LSTM,64,128,1,170.77305080579674
GRU,16,32,1,23.525951261105746
GRU,16,64,1,20.953380501788594
GRU,16,128,1,50.10249718375828
GRU,32,32,1,149.08109225397524
GRU,32,64,1,40.437404922817066
GRU,32,128,1,34.31250389762547
GRU,64,32,1,154.5310813240383
GRU,64,64,1,47.95653426128885
GRU,64,128,1,263.4846937760063
RNN,16,32,5,326.2015085634978
RNN,16,64,5,384.77328723409903
RNN,16,128,5,1099.3475819463315
RNN,32,32,5,246.47076250159222
RNN,32,64,5,329.85308439835256
RNN,32,128,5,1098.7943765391474
RNN,64,32,5,477.369012583857
RNN,64,64,5,365.6013999607252
RNN,64,128,5,413.95318603515625
LSTM,16,32,5,24.514795303344727
LSTM,16,64,5,24.9299375285273
LSTM,16,128,5,34.382882035296895
LSTM,32,32,5,171.14435179337212
LSTM,32,64,5,185.21222504325536
LSTM,32,128,5,24.800554503565248
LSTM,64,32,5,183.13448831309444
LSTM,64,64,5,37.47129328354545
LSTM,64,128,5,51.097612007804535
GRU,16,32,5,30.631973432457965
GRU,16,64,5,30.393669625987176
GRU,16,128,5,55.05557930987814
GRU,32,32,5,54.14869615306025
GRU,32,64,5,167.3384770932405
GRU,32,128,5,50.49319474593453
GRU,64,32,5,170.90516430398694
GRU,64,64,5,153.55044091266134
GRU,64,128,5,53.40392129317574
RNN,16,32,10,285.0945102857507
RNN,16,64,10,361.18107339610225
RNN,16,128,10,320.00166884712553
RNN,32,32,10,359.23063095756197
RNN,32,64,10,351.7510011092476
RNN,32,128,10,126.81333442356275
RNN,64,32,10,336.1739501953125
RNN,64,64,10,342.61622918170434
RNN,64,128,10,208.79256007982337
LSTM,16,32,10,47.18654777692712
LSTM,16,64,10,35.0142301061879
LSTM,16,128,10,31.15265079166578
LSTM,32,32,10,57.50707398290219
LSTM,32,64,10,21.263865159905475
LSTM,32,128,10,39.93263937079388
LSTM,64,32,10,57.33804578366487
LSTM,64,64,10,52.7249521587206
LSTM,64,128,10,59.768143446549125
GRU,16,32,10,28.489941555520762
GRU,16,64,10,25.493946241295856
GRU,16,128,10,35.874996682871945
GRU,32,32,10,39.71779649154
GRU,32,64,10,33.58222509467083
GRU,32,128,10,36.87489372750987
GRU,64,32,10,55.22760449285092
GRU,64,64,10,168.17254008417544
GRU,64,128,10,76.96743749535602
RNN,16,32,30,245.38945504893428
RNN,16,64,30,303.7835593845533
RNN,16,128,30,1168.2417987325916
RNN,32,32,30,221.63304668924084
RNN,32,64,30,274.024205746858
RNN,32,128,30,1167.9884391452956
RNN,64,32,30,316.6191160782524
RNN,64,64,30,252.00204235574475
RNN,64,128,30,1168.0314676036005
LSTM,16,32,30,46.3279622119406
LSTM,16,64,30,32.77271196116572
LSTM,16,128,30,1168.1180778171706
LSTM,32,32,30,192.53192719169286
LSTM,32,64,30,40.278250404026195
LSTM,32,128,30,35.424202234848686
LSTM,64,32,30,51.08382113083549
LSTM,64,64,30,186.7356419770614
LSTM,64,128,30,177.50895442133364
GRU,16,32,30,35.31677892933721
GRU,16,64,30,34.573822394661285
GRU,16,128,30,55.778304846390434
GRU,32,32,30,65.31020172782566
GRU,32,64,30,37.21405522719674
GRU,32,128,30,38.621485295503035
GRU,64,32,30,120.50829563970152
GRU,64,64,30,75.70721891651984
GRU,64,128,30,189.14709671683934
1 Model Hidden Size RNN Size Step Final Loss
2 RNN 16 32 1 372.1873533829399
3 RNN 16 64 1 371.67606353759766
4 RNN 16 128 1 328.19922903309697
5 RNN 32 32 1 376.71770842179006
6 RNN 32 64 1 369.6340640524159
7 RNN 32 128 1 281.8589802617612
8 RNN 64 32 1 378.75125188412875
9 RNN 64 64 1 405.08673095703125
10 RNN 64 128 1 1088.4633191979449
11 LSTM 16 32 1 152.01849580847698
12 LSTM 16 64 1 152.46532937754756
13 LSTM 16 128 1 39.81503051260243
14 LSTM 32 32 1 45.500295265861176
15 LSTM 32 64 1 51.218478575996734
16 LSTM 32 128 1 32.29804623645285
17 LSTM 64 32 1 105.44863045733908
18 LSTM 64 64 1 62.87264318051545
19 LSTM 64 128 1 170.77305080579674
20 GRU 16 32 1 23.525951261105746
21 GRU 16 64 1 20.953380501788594
22 GRU 16 128 1 50.10249718375828
23 GRU 32 32 1 149.08109225397524
24 GRU 32 64 1 40.437404922817066
25 GRU 32 128 1 34.31250389762547
26 GRU 64 32 1 154.5310813240383
27 GRU 64 64 1 47.95653426128885
28 GRU 64 128 1 263.4846937760063
29 RNN 16 32 5 326.2015085634978
30 RNN 16 64 5 384.77328723409903
31 RNN 16 128 5 1099.3475819463315
32 RNN 32 32 5 246.47076250159222
33 RNN 32 64 5 329.85308439835256
34 RNN 32 128 5 1098.7943765391474
35 RNN 64 32 5 477.369012583857
36 RNN 64 64 5 365.6013999607252
37 RNN 64 128 5 413.95318603515625
38 LSTM 16 32 5 24.514795303344727
39 LSTM 16 64 5 24.9299375285273
40 LSTM 16 128 5 34.382882035296895
41 LSTM 32 32 5 171.14435179337212
42 LSTM 32 64 5 185.21222504325536
43 LSTM 32 128 5 24.800554503565248
44 LSTM 64 32 5 183.13448831309444
45 LSTM 64 64 5 37.47129328354545
46 LSTM 64 128 5 51.097612007804535
47 GRU 16 32 5 30.631973432457965
48 GRU 16 64 5 30.393669625987176
49 GRU 16 128 5 55.05557930987814
50 GRU 32 32 5 54.14869615306025
51 GRU 32 64 5 167.3384770932405
52 GRU 32 128 5 50.49319474593453
53 GRU 64 32 5 170.90516430398694
54 GRU 64 64 5 153.55044091266134
55 GRU 64 128 5 53.40392129317574
56 RNN 16 32 10 285.0945102857507
57 RNN 16 64 10 361.18107339610225
58 RNN 16 128 10 320.00166884712553
59 RNN 32 32 10 359.23063095756197
60 RNN 32 64 10 351.7510011092476
61 RNN 32 128 10 126.81333442356275
62 RNN 64 32 10 336.1739501953125
63 RNN 64 64 10 342.61622918170434
64 RNN 64 128 10 208.79256007982337
65 LSTM 16 32 10 47.18654777692712
66 LSTM 16 64 10 35.0142301061879
67 LSTM 16 128 10 31.15265079166578
68 LSTM 32 32 10 57.50707398290219
69 LSTM 32 64 10 21.263865159905475
70 LSTM 32 128 10 39.93263937079388
71 LSTM 64 32 10 57.33804578366487
72 LSTM 64 64 10 52.7249521587206
73 LSTM 64 128 10 59.768143446549125
74 GRU 16 32 10 28.489941555520762
75 GRU 16 64 10 25.493946241295856
76 GRU 16 128 10 35.874996682871945
77 GRU 32 32 10 39.71779649154
78 GRU 32 64 10 33.58222509467083
79 GRU 32 128 10 36.87489372750987
80 GRU 64 32 10 55.22760449285092
81 GRU 64 64 10 168.17254008417544
82 GRU 64 128 10 76.96743749535602
83 RNN 16 32 30 245.38945504893428
84 RNN 16 64 30 303.7835593845533
85 RNN 16 128 30 1168.2417987325916
86 RNN 32 32 30 221.63304668924084
87 RNN 32 64 30 274.024205746858
88 RNN 32 128 30 1167.9884391452956
89 RNN 64 32 30 316.6191160782524
90 RNN 64 64 30 252.00204235574475
91 RNN 64 128 30 1168.0314676036005
92 LSTM 16 32 30 46.3279622119406
93 LSTM 16 64 30 32.77271196116572
94 LSTM 16 128 30 1168.1180778171706
95 LSTM 32 32 30 192.53192719169286
96 LSTM 32 64 30 40.278250404026195
97 LSTM 32 128 30 35.424202234848686
98 LSTM 64 32 30 51.08382113083549
99 LSTM 64 64 30 186.7356419770614
100 LSTM 64 128 30 177.50895442133364
101 GRU 16 32 30 35.31677892933721
102 GRU 16 64 30 34.573822394661285
103 GRU 16 128 30 55.778304846390434
104 GRU 32 32 30 65.31020172782566
105 GRU 32 64 30 37.21405522719674
106 GRU 32 128 30 38.621485295503035
107 GRU 64 32 30 120.50829563970152
108 GRU 64 64 30 75.70721891651984
109 GRU 64 128 30 189.14709671683934
+38 -16
View File
@@ -126,31 +126,53 @@ class BaseDataset(Dataset):
torch.save(data_dict, filepath)
class SaveDataset(Dataset):
def __init__(self, data_dict, device="cpu", step=1):
def __init__(self, data_dict, step=1, start_offset=0, end_offset=0):
super().__init__()
self.X_feat = data_dict["X_feat"].to(device)
self.X_time = data_dict["X_time"].to(device)
self.Y_out = data_dict["Y_out"].to(device)
self.step = step
self.device = device
if "X_context" in data_dict:
self.X_context = data_dict["X_context"].to(device)
else:
self.X_context = None
self.start_offset = start_offset // step
self.end_offset = end_offset // step
# Keep everything **on CPU**
self.X_feat = data_dict["X_feat"]
self.X_time = data_dict["X_time"]
self.Y_out = data_dict["Y_out"]
self.X_context = data_dict.get("X_context", None)
# Precompute slice indices (avoids Python overhead in worker processes)
max_len = self.X_time.shape[1]
t0 = self.start_offset
t1 = max_len - self.end_offset
self.idx_feat = torch.arange(t0, self.X_feat.shape[1], self.step)
self.idx_time = torch.arange(t0, t1, self.step)
self.idx_out = torch.arange(0, self.Y_out.shape[1] - self.end_offset, self.step)
def __len__(self):
return self.X_feat.shape[0]
def __getitem__(self, idx):
X_f = self.X_feat[idx,::self.step]
X_t = self.X_time[idx,::self.step]
Y = self.Y_out[idx,::self.step]
if self.X_context is not None:
X_c = self.X_context[idx]
else:
X_c = None
X_f = self.X_feat[idx].index_select(0, self.idx_feat)
X_t = self.X_time[idx].index_select(0, self.idx_time)
Y = self.Y_out[idx].index_select(0, self.idx_out)
X_c = self.X_context[idx] if self.X_context is not None else None
return X_f, X_t, Y, X_c
def collate_to_cuda(device):
def _collate(batch):
Xf, Xt, Y, Xc = zip(*batch)
Xf = torch.stack(Xf).to(device)
Xt = torch.stack(Xt).to(device)
Y = torch.stack(Y).to(device)
if Xc[0] is not None:
Xc = torch.stack(Xc).to(device)
else:
Xc = None
return Xf, Xt, Y, Xc
return _collate
class EvenlySpacedDataset(BaseDataset):
def __init__(
self,
+38 -23
View File
@@ -4,16 +4,24 @@ import torch.nn as nn
class BaseRNN(nn.Module):
def __init__(self, time_in, feat_in, context_in, hidden_size,
rnn_size, out_size, rnn_type="RNN", device="cpu"):
rnn_size, out_size, hidden_layers=1, rnn_layers=3, rnn_dropout=0.0, rnn_type="RNN", device="cpu"):
super().__init__()
base_unit = lambda in_size, out_size: nn.Sequential(
nn.Linear(in_size, out_size, device=device),
nn.GELU(),
nn.Linear(out_size, out_size, device=device),
nn.GELU(),
nn.LayerNorm(out_size, device=device)
)
def base_unit(in_size, out_size):
layers = []
# First projection
layers.append(nn.Linear(in_size, out_size, device=device))
layers.append(nn.GELU())
# Hidden repeated blocks: (Linear → GELU) * hidden_layers
for _ in range(hidden_layers):
layers.append(nn.Linear(out_size, out_size, device=device))
layers.append(nn.GELU())
# Final normalization
layers.append(nn.LayerNorm(out_size, device=device))
return nn.Sequential(*layers)
self.time_proj = base_unit(time_in, hidden_size)
self.feat_proj = base_unit(feat_in, hidden_size)
@@ -25,24 +33,27 @@ class BaseRNN(nn.Module):
self.rnn = nn.RNN(
input_size=hidden_size*3,
hidden_size=rnn_size,
num_layers=3,
num_layers=rnn_layers,
batch_first=True,
dropout=rnn_dropout,
device=device
)
elif rnn_type == "LSTM":
self.rnn = nn.LSTM(
input_size=hidden_size*3,
hidden_size=rnn_size,
num_layers=3,
num_layers=rnn_layers,
batch_first=True,
dropout=rnn_dropout,
device=device
)
elif rnn_type == "GRU":
self.rnn = nn.GRU(
input_size=hidden_size*3,
hidden_size=rnn_size,
num_layers=3,
num_layers=rnn_layers,
batch_first=True,
dropout=rnn_dropout,
device=device
)
else:
@@ -101,8 +112,8 @@ class TwoInputRNN(BaseRNN):
After init_steps, only time inputs are provided and the model predicts a sequence.
Predictions can be compared to targets with a specified `offset`.
"""
def __init__(self, time_in, feat_in, hidden_size, rnn_size, out_size, device="cpu"):
super().__init__(time_in, feat_in, context_in=None, hidden_size=hidden_size, rnn_size=rnn_size, out_size=out_size, device=device)
def __init__(self, time_in, feat_in, hidden_size, rnn_size, out_size, hidden_layers=1, rnn_layers=3, rnn_dropout=0.0, device="cpu"):
super().__init__(time_in, feat_in, context_in=None, hidden_size=hidden_size, rnn_size=rnn_size, out_size=out_size, hidden_layers=hidden_layers, rnn_layers=rnn_layers, rnn_dropout=rnn_dropout, device=device)
class ThreeInputRNN(BaseRNN):
"""
@@ -113,8 +124,8 @@ class ThreeInputRNN(BaseRNN):
After init_steps, only time and context inputs are provided and the model predicts a sequence.
Predictions can be compared to targets with a specified `offset`.
"""
def __init__(self, time_in, feat_in, context_in, hidden_size, rnn_size, out_size, device="cpu"):
super().__init__(time_in, feat_in, context_in=context_in, hidden_size=hidden_size, rnn_size=rnn_size, out_size=out_size, device=device)
def __init__(self, time_in, feat_in, context_in, hidden_size, rnn_size, out_size, hidden_layers=1, rnn_layers=3, rnn_dropout=0.0, device="cpu"):
super().__init__(time_in, feat_in, context_in=context_in, hidden_size=hidden_size, rnn_size=rnn_size, out_size=out_size, hidden_layers=hidden_layers, rnn_layers=rnn_layers, rnn_dropout=rnn_dropout, device=device)
class TwoInputLSTM(BaseRNN):
"""
@@ -124,8 +135,8 @@ class TwoInputLSTM(BaseRNN):
After init_steps, only time inputs are provided and the model predicts a sequence.
Predictions can be compared to targets with a specified `offset`.
"""
def __init__(self, time_in, feat_in, hidden_size, rnn_size, out_size, device="cpu"):
super().__init__(time_in, feat_in, context_in=None, hidden_size=hidden_size, rnn_size=rnn_size, out_size=out_size, rnn_type="LSTM", device=device)
def __init__(self, time_in, feat_in, hidden_size, rnn_size, out_size, hidden_layers=1, rnn_layers=3, rnn_dropout=0.0, device="cpu"):
super().__init__(time_in, feat_in, context_in=None, hidden_size=hidden_size, rnn_size=rnn_size, out_size=out_size, hidden_layers=hidden_layers, rnn_layers=rnn_layers, rnn_dropout=rnn_dropout, rnn_type="LSTM", device=device)
class ThreeInputLSTM(BaseRNN):
"""
@@ -136,8 +147,8 @@ class ThreeInputLSTM(BaseRNN):
After init_steps, only time and context inputs are provided and the model predicts a sequence.
Predictions can be compared to targets with a specified `offset`.
"""
def __init__(self, time_in, feat_in, context_in, hidden_size, rnn_size, out_size, device="cpu"):
super().__init__(time_in, feat_in, context_in=context_in, hidden_size=hidden_size, rnn_size=rnn_size, out_size=out_size, rnn_type="LSTM", device=device)
def __init__(self, time_in, feat_in, context_in, hidden_size, rnn_size, out_size, hidden_layers=1, rnn_layers=3, rnn_dropout=0.0, device="cpu"):
super().__init__(time_in, feat_in, context_in=context_in, hidden_size=hidden_size, rnn_size=rnn_size, out_size=out_size, hidden_layers=hidden_layers, rnn_layers=rnn_layers, rnn_dropout=rnn_dropout, rnn_type="LSTM", device=device)
class TwoInputGRU(BaseRNN):
"""
@@ -147,8 +158,8 @@ class TwoInputGRU(BaseRNN):
After init_steps, only time inputs are provided and the model predicts a sequence.
Predictions can be compared to targets with a specified `offset`.
"""
def __init__(self, time_in, feat_in, hidden_size, rnn_size, out_size, device="cpu"):
super().__init__(time_in, feat_in, context_in=None, hidden_size=hidden_size, rnn_size=rnn_size, out_size=out_size, rnn_type="GRU", device=device)
def __init__(self, time_in, feat_in, hidden_size, rnn_size, out_size, hidden_layers=1, rnn_layers=3, rnn_dropout=0.0, device="cpu"):
super().__init__(time_in, feat_in, context_in=None, hidden_size=hidden_size, rnn_size=rnn_size, out_size=out_size, hidden_layers=hidden_layers, rnn_layers=rnn_layers, rnn_dropout=rnn_dropout, rnn_type="GRU", device=device)
class ThreeInputGRU(BaseRNN):
"""
@@ -159,5 +170,9 @@ class ThreeInputGRU(BaseRNN):
After init_steps, only time and context inputs are provided and the model predicts a sequence.
Predictions can be compared to targets with a specified `offset`.
"""
def __init__(self, time_in, feat_in, context_in, hidden_size, rnn_size, out_size, device="cpu"):
super().__init__(time_in, feat_in, context_in=context_in, hidden_size=hidden_size, rnn_size=rnn_size, out_size=out_size, rnn_type="GRU", device=device)
def __init__(self, time_in, feat_in, context_in, hidden_size, rnn_size, out_size, hidden_layers=1, rnn_layers=3, rnn_dropout=0.0, device="cpu"):
super().__init__(time_in, feat_in, context_in=context_in, hidden_size=hidden_size, rnn_size=rnn_size, out_size=out_size, hidden_layers=hidden_layers, rnn_layers=rnn_layers, rnn_dropout=rnn_dropout, rnn_type="GRU", device=device)
OptimalRNN_cuda = ThreeInputGRU(time_in=2, feat_in=4, context_in=5, hidden_size=32, rnn_size=64, out_size=3, device="cuda")
OptimalRNN_cpu = ThreeInputGRU(time_in=2, feat_in=4, context_in=5, hidden_size=32, rnn_size=64, out_size=3, device="cpu")