upper bound get_lags_for_frequency by context_length

This commit is contained in:
Dr. Kashif Rasul
2021-03-16 16:51:57 +01:00
parent 3a0223d171
commit b01e7c6a24
2 changed files with 3 additions and 6 deletions
+3 -1
View File
@@ -91,7 +91,9 @@ class DeepAREstimator(PyTorchEstimator):
)
self.scaling = scaling
self.lags_seq = (
lags_seq if lags_seq is not None else get_lags_for_frequency(freq_str=freq)
lags_seq
if lags_seq is not None
else get_lags_for_frequency(freq_str=freq, lag_ub=self.context_length)
)
self.time_features = (
time_features
@@ -9,11 +9,6 @@ from gluonts.torch.model.predictor import PyTorchPredictor
from gluonts.torch.modules.distribution_output import DistributionOutput
from gluonts.model.predictor import Predictor
from gluonts.dataset.field_names import FieldName
from gluonts.time_feature import (
TimeFeature,
get_lags_for_frequency,
time_features_from_frequency_str,
)
from gluonts.transform import (
Transformation,
Chain,