Add more hyperparameter scan modes
This commit is contained in:
@@ -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