Add more hyperparameter scan modes
This commit is contained in:
File diff suppressed because one or more lines are too long
@@ -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
|
||||
|
@@ -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,
|
||||
|
||||
@@ -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")
|
||||
Reference in New Issue
Block a user