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(