mirror of
https://github.com/wassname/pytorch-ts.git
synced 2026-07-22 13:00:05 +08:00
63 KiB
63 KiB
In [17]:
import matplotlib.pyplot as plt
import jsonIn [2]:
import torchIn [3]:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")In [4]:
from pts.dataset.repository import get_dataset
from pts.dataset.utils import to_pandasIn [5]:
dataset = get_dataset("m5", regenerate=False)In [25]:
entry = next(iter(dataset.train))
train_series = to_pandas(entry)
train_series.plot()
plt.grid(which="both")
plt.legend(["train series"], loc="upper left")
plt.title(entry['item_id'])
plt.show()In [24]:
entry = next(iter(dataset.test))
test_series = to_pandas(entry)
test_series.plot()
plt.axvline(train_series.index[-1], color='r') # end of train dataset
plt.grid(which="both")
plt.legend(["test series", "end of train series"], loc="upper left")
plt.title(entry['item_id'])
plt.show()In [8]:
print(f"Recommended prediction horizon: {dataset.metadata.prediction_length}")
print(f"Frequency of the time series: {dataset.metadata.freq}")Recommended prediction horizon: 28 Frequency of the time series: D
In [9]:
from pts.model.deepar import DeepAREstimator
from pts.modules import ZeroInflatedNegativeBinomialOutput
from pts import TrainerIn [10]:
estimator = DeepAREstimator(
distr_output=ZeroInflatedNegativeBinomialOutput(),
cell_type='GRU',
input_size=72,
num_cells=64,
num_layers=3,
dropout_rate=0.2,
use_feat_dynamic_real=True,
use_feat_static_cat=True,
cardinality=[int(cat_feat_info.cardinality) for cat_feat_info in dataset.metadata.feat_static_cat],
embedding_dimension = [4, 4, 4, 4, 16],
prediction_length=dataset.metadata.prediction_length,
context_length=dataset.metadata.prediction_length*2,
freq=dataset.metadata.freq,
scaling=True,
trainer=Trainer(device=device,
epochs=50,
learning_rate=1e-3,
num_batches_per_epoch=120,
batch_size=256,
num_workers=8,
pin_memory=True,
)
)In [11]:
predictor = estimator.train(dataset.train)119it [00:29, 4.09it/s, avg_epoch_loss=1.17, epoch=0] 119it [00:31, 3.81it/s, avg_epoch_loss=1.14, epoch=1] 119it [00:27, 4.29it/s, avg_epoch_loss=1.12, epoch=2] 119it [00:29, 4.06it/s, avg_epoch_loss=1.11, epoch=3] 119it [00:28, 4.17it/s, avg_epoch_loss=1.1, epoch=4] 119it [00:28, 4.14it/s, avg_epoch_loss=1.1, epoch=5] 119it [00:30, 3.93it/s, avg_epoch_loss=1.1, epoch=6] 119it [00:27, 4.27it/s, avg_epoch_loss=1.11, epoch=7] 119it [00:28, 4.16it/s, avg_epoch_loss=1.09, epoch=8] 119it [00:29, 4.09it/s, avg_epoch_loss=1.11, epoch=9] 119it [00:27, 4.40it/s, avg_epoch_loss=1.1, epoch=10] 119it [00:28, 4.23it/s, avg_epoch_loss=1.1, epoch=11] 119it [00:28, 4.12it/s, avg_epoch_loss=1.1, epoch=12] 119it [00:29, 4.06it/s, avg_epoch_loss=1.1, epoch=13] 119it [00:29, 4.10it/s, avg_epoch_loss=1.11, epoch=14] 119it [00:28, 4.24it/s, avg_epoch_loss=1.1, epoch=15] 119it [00:30, 3.95it/s, avg_epoch_loss=1.1, epoch=16] 119it [00:27, 4.28it/s, avg_epoch_loss=1.09, epoch=17] 119it [00:27, 4.26it/s, avg_epoch_loss=1.1, epoch=18] 119it [00:29, 4.07it/s, avg_epoch_loss=1.1, epoch=19] 119it [00:29, 3.98it/s, avg_epoch_loss=1.09, epoch=20] 119it [00:27, 4.33it/s, avg_epoch_loss=1.1, epoch=21] 119it [00:29, 4.08it/s, avg_epoch_loss=1.09, epoch=22] 119it [00:29, 4.09it/s, avg_epoch_loss=1.09, epoch=23] 119it [00:28, 4.22it/s, avg_epoch_loss=1.09, epoch=24] 119it [00:27, 4.26it/s, avg_epoch_loss=1.09, epoch=25] 119it [00:31, 3.81it/s, avg_epoch_loss=1.1, epoch=26] 119it [00:31, 3.73it/s, avg_epoch_loss=1.09, epoch=27] 119it [00:27, 4.32it/s, avg_epoch_loss=1.08, epoch=28] 119it [00:28, 4.14it/s, avg_epoch_loss=1.09, epoch=29] 119it [00:30, 3.87it/s, avg_epoch_loss=1.08, epoch=30] 119it [00:28, 4.19it/s, avg_epoch_loss=1.09, epoch=31] 119it [00:28, 4.17it/s, avg_epoch_loss=1.08, epoch=32] 119it [00:29, 4.09it/s, avg_epoch_loss=1.1, epoch=33] 119it [00:27, 4.39it/s, avg_epoch_loss=1.09, epoch=34] 119it [00:28, 4.21it/s, avg_epoch_loss=1.09, epoch=35] 119it [00:28, 4.16it/s, avg_epoch_loss=1.09, epoch=36] 119it [00:27, 4.31it/s, avg_epoch_loss=1.08, epoch=37] 119it [00:29, 4.07it/s, avg_epoch_loss=1.09, epoch=38] 119it [00:28, 4.19it/s, avg_epoch_loss=1.09, epoch=39] 119it [00:29, 4.06it/s, avg_epoch_loss=1.09, epoch=40] 119it [00:28, 4.14it/s, avg_epoch_loss=1.08, epoch=41] 119it [00:28, 4.16it/s, avg_epoch_loss=1.09, epoch=42] 119it [00:27, 4.25it/s, avg_epoch_loss=1.09, epoch=43] 119it [00:27, 4.26it/s, avg_epoch_loss=1.1, epoch=44] 119it [00:26, 4.41it/s, avg_epoch_loss=1.09, epoch=45] 119it [00:27, 4.25it/s, avg_epoch_loss=1.08, epoch=46] 119it [00:28, 4.20it/s, avg_epoch_loss=1.09, epoch=47] 119it [00:30, 3.92it/s, avg_epoch_loss=1.09, epoch=48] 119it [00:30, 3.96it/s, avg_epoch_loss=1.09, epoch=49]
In [12]:
from pts.evaluation import make_evaluation_predictions, EvaluatorIn [13]:
forecast_it, ts_it = make_evaluation_predictions(
dataset=dataset.test, # test dataset
predictor=predictor, # predictor
num_samples=100, # number of sample paths we want for evaluation
)In [14]:
forecasts = list(forecast_it)
tss = list(ts_it)In [15]:
evaluator = Evaluator()
agg_metrics, item_metrics = evaluator(iter(tss), iter(forecasts), num_series=len(dataset.test))Running evaluation: 100%|██████████| 30490/30490 [00:01<00:00, 23152.64it/s]
In [18]:
print(json.dumps(agg_metrics, indent=4)){
"MSE": 4.439620313754265,
"abs_error": 807089.0,
"abs_target_sum": 1231764.0,
"abs_target_mean": 1.4428196598416343,
"seasonal_error": 1.1272178349378457,
"MASE": 0.8789472000957106,
"MAPE": 0.30587227335898637,
"sMAPE": 0.6816909686747539,
"OWA": NaN,
"MSIS": 7.28829495943688,
"QuantileLoss[0.1]": 228315.8,
"Coverage[0.1]": 0.0042929766199690765,
"QuantileLoss[0.2]": 422650.8,
"Coverage[0.2]": 0.01732535257461463,
"QuantileLoss[0.3]": 586642.4,
"Coverage[0.3]": 0.042479970013587595,
"QuantileLoss[0.4]": 716891.6,
"Coverage[0.4]": 0.08317012603663966,
"QuantileLoss[0.5]": 807089.0,
"Coverage[0.5]": 0.14288174108607035,
"QuantileLoss[0.6]": 854345.2,
"Coverage[0.6]": 0.2176439582064377,
"QuantileLoss[0.7]": 842037.0,
"Coverage[0.7]": 0.3306306517359322,
"QuantileLoss[0.8]": 755156.7999999999,
"Coverage[0.8]": 0.48669352949444783,
"QuantileLoss[0.9]": 547328.7999999999,
"Coverage[0.9]": 0.7026999484608537,
"RMSE": 2.1070406530853325,
"NRMSE": 1.4603631429007586,
"ND": 0.655230222672525,
"wQuantileLoss[0.1]": 0.185356772888313,
"wQuantileLoss[0.2]": 0.34312644305240286,
"wQuantileLoss[0.3]": 0.47626201122942385,
"wQuantileLoss[0.4]": 0.5820040202506324,
"wQuantileLoss[0.5]": 0.655230222672525,
"wQuantileLoss[0.6]": 0.6935948769407126,
"wQuantileLoss[0.7]": 0.6836025407464417,
"wQuantileLoss[0.8]": 0.6130693866682253,
"wQuantileLoss[0.9]": 0.44434550774336634,
"mean_wQuantileLoss": 0.5196213091324493,
"MAE_Coverage": 0.2746868606412719
}
In [19]:
item_metrics.plot(x='MSIS', y='MASE', kind='scatter')
plt.grid(which="both")
plt.show()In [ ]: