From a1fe6832547e0f1056afa6f9eed52fdc43519363 Mon Sep 17 00:00:00 2001 From: "Dr. Kashif Rasul" Date: Tue, 14 Jan 2020 11:55:33 +0100 Subject: [PATCH] remove comments and cleanup --- pts/model/tempflow/tempflow_network.py | 6 +----- 1 file changed, 1 insertion(+), 5 deletions(-) diff --git a/pts/model/tempflow/tempflow_network.py b/pts/model/tempflow/tempflow_network.py index 9e30b68..5f1df6b 100644 --- a/pts/model/tempflow/tempflow_network.py +++ b/pts/model/tempflow/tempflow_network.py @@ -267,8 +267,6 @@ class TempFlowTrainingNetwork(nn.Module): past_observed_values[:, -self.context_length :, ...], ) - # import pdb; pdb.set_trace() - outputs, states, lags_scaled, inputs = self.unroll( lags=lags, scale=scale, @@ -363,7 +361,7 @@ class TempFlowTrainingNetwork(nn.Module): # unroll the decoder in "training mode", i.e. by providing future data # as well - rnn_outputs, _, scale, _, inputs = self.unroll_encoder( + rnn_outputs, _, scale, _, _ = self.unroll_encoder( past_time_feat=past_time_feat, past_target_cdf=past_target_cdf, past_observed_values=past_observed_values, @@ -387,8 +385,6 @@ class TempFlowTrainingNetwork(nn.Module): # (batch_size, subseq_length, 1) likelihoods = -distr.log_prob(target).unsqueeze(-1) - # import pdb; pdb.set_trace() - # assert_shape(likelihoods, (-1, seq_len, 1)) past_observed_values = torch.min(