diff --git a/Code/python/src/aiRNN/losses.py b/Code/python/src/aiRNN/losses.py index 23cf738..b211e39 100644 --- a/Code/python/src/aiRNN/losses.py +++ b/Code/python/src/aiRNN/losses.py @@ -6,7 +6,7 @@ import math class HaversineMSEAltitudeLoss(nn.Module): - def __init__(self, alt_const=1.0, earth_radius_km=6371.0): + def __init__(self, alt_const=1e-4, earth_radius_km=6371.0): super().__init__() self.alt_const = alt_const self.R = earth_radius_km @@ -47,7 +47,7 @@ class HaversineMSEAltitudeLoss(nn.Module): return haversine_mse + alt_penalty class HaversineAltitudeLoss(nn.Module): - def __init__(self, alt_const=1.0, earth_radius_km=6371.0): + def __init__(self, alt_const=1e-4, earth_radius_km=6371.0): super().__init__() self.alt_const = alt_const self.R = earth_radius_km diff --git a/Code/python/src/aiRNN/models.py b/Code/python/src/aiRNN/models.py index 891ffbb..9dfc0d4 100644 --- a/Code/python/src/aiRNN/models.py +++ b/Code/python/src/aiRNN/models.py @@ -7,17 +7,24 @@ class BaseRNN(nn.Module): rnn_size, out_size, rnn_type="RNN", device="cpu"): super().__init__() - self.time_proj = nn.Linear(time_in, hidden_size, device=device) - self.feat_proj = nn.Linear(feat_in, hidden_size, device=device) + base_unit = lambda in_size, out_size: nn.Sequential( + nn.Linear(in_size, out_size, device=device), + nn.ReLU(), + nn.Linear(out_size, out_size, device=device), + nn.ReLU(), + ) + + self.time_proj = base_unit(time_in, hidden_size) + self.feat_proj = base_unit(feat_in, hidden_size) self.context_proj = ( - nn.Linear(context_in, hidden_size, device=device) if context_in is not None else None + base_unit(context_in, hidden_size) if context_in is not None else None ) if rnn_type == "RNN": self.rnn = nn.RNN( input_size=hidden_size*3, hidden_size=rnn_size, - num_layers=2, + num_layers=3, batch_first=True, device=device ) @@ -25,7 +32,7 @@ class BaseRNN(nn.Module): self.rnn = nn.LSTM( input_size=hidden_size*3, hidden_size=rnn_size, - num_layers=2, + num_layers=3, batch_first=True, device=device ) @@ -33,8 +40,11 @@ class BaseRNN(nn.Module): raise ValueError("Unsupported rnn_type") readout_in = rnn_size - self.readout = nn.Linear(readout_in, out_size, device=device) - + self.readout = nn.Sequential( + nn.Linear(readout_in, hidden_size, device=device), + nn.ReLU(), + nn.Linear(hidden_size, out_size, device=device), + ) self.context_in = context_in self.rnn_size = rnn_size