This commit is contained in:
Dr. Kashif Rasul
2020-01-04 12:41:21 +01:00
parent 71fff3c411
commit c1be3fa5a7
3 changed files with 7 additions and 7 deletions
+1 -1
View File
@@ -283,7 +283,7 @@ class DeepARTrainingNetwork(DeepARNetwork):
else observed_values.min(dim=-1, keepdim=False)
)
weighted_loss = weighted_average(loss, loss_weights)
weighted_loss = weighted_average(loss, weights=loss_weights)
return weighted_loss, loss
+5 -5
View File
@@ -153,7 +153,7 @@ class DeepVARTrainingNetwork(nn.Module):
# (batch_size, seq_len, target_dim * embed_dim)
repeated_index_embeddings = (
index_embeddings.unsqueeze(1)
.expand(-1, unroll_length, -1)
.expand(-1, unroll_length, -1, -1)
.reshape((-1, unroll_length, self.target_dim * self.embed_dim))
)
@@ -263,7 +263,7 @@ class DeepVARTrainingNetwork(nn.Module):
# scale shape is (batch_size, 1, target_dim)
_, scale = self.scaler(
past_target_cdf[:, -self.context_length :, ...],
past_observed_values[:, -self.context_length : ...,],
past_observed_values[:, -self.context_length :, ...],
)
outputs, states, lags_scaled, inputs = self.unroll(
@@ -401,17 +401,17 @@ class DeepVARTrainingNetwork(nn.Module):
# mask the loss at one time step if one or more observations is missing
# in the target dimensions (batch_size, subseq_length, 1)
loss_weights = observed_values.min(dim=-1, keepdim=True)
loss_weights,_ = observed_values.min(dim=-1, keepdim=True)
# assert_shape(loss_weights, (-1, seq_len, 1))
loss = weighted_average(x=likelihoods, weights=loss_weights, dim=1)
loss = weighted_average(likelihoods, weights=loss_weights, dim=1)
# assert_shape(loss, (-1, -1, 1))
self.distribution = distr
return (loss, likelihoods) + distr_args
return (loss.sum(), likelihoods) + distr_args
class DeepVARPredictionNetwork(DeepVARTrainingNetwork):
+1 -1
View File
@@ -87,7 +87,7 @@ def test_deepvar(
):
estimator = Estimator(
input_size=10,
input_size=44,
num_cells=20,
num_layers=1,
pick_incomplete=True,