mirror of
https://github.com/wassname/pytorch-ts.git
synced 2026-08-11 11:24:31 +08:00
Gluon master (#29)
* Estimator needs an create_instance_splitter now * updated estimators and tests * fix test * validated
This commit is contained in:
committed by
GitHub Enterprise
parent
d5cef439af
commit
ea9b2b7df5
@@ -3,10 +3,7 @@ from typing import List
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.distributions import (
|
||||
Normal,
|
||||
Uniform,
|
||||
Bernoulli)
|
||||
from torch.distributions import Normal, Uniform, Bernoulli
|
||||
from torch.nn.utils import clip_grad_norm_
|
||||
from torch.optim import SGD
|
||||
from torch.utils.data import TensorDataset, DataLoader
|
||||
@@ -17,10 +14,7 @@ from gluonts.evaluation.backtest import make_evaluation_predictions
|
||||
from gluonts.torch.modules.distribution_output import DistributionOutput
|
||||
from pts import Trainer
|
||||
from pts.model.deepar import DeepAREstimator
|
||||
from pts.model.simple_feedforward import SimpleFeedForwardEstimator
|
||||
from pts.modules import (
|
||||
ImplicitQuantileOutput
|
||||
)
|
||||
from pts.modules import ImplicitQuantileOutput
|
||||
|
||||
NUM_SAMPLES = 2000
|
||||
BATCH_SIZE = 32
|
||||
@@ -73,11 +67,15 @@ def learn_distribution(
|
||||
i, (data, sample_label) = next(enumerate(sampling_dataloader))
|
||||
distr_args = arg_proj(data)
|
||||
distr = distr_output.distribution(distr_args)
|
||||
samples = distr.sample((NUM_SAMPLES, ))
|
||||
samples = distr.sample((NUM_SAMPLES,))
|
||||
|
||||
with torch.no_grad():
|
||||
percentile_90 = distr.quantile_function(torch.ones((1, 1, 1)), torch.ones((1, 1)) * 0.9)
|
||||
percentile_10 = distr.quantile_function(torch.ones((1, 1, 1)), torch.ones((1, 1)) * 0.1)
|
||||
percentile_90 = distr.quantile_function(
|
||||
torch.ones((1, 1, 1)), torch.ones((1, 1)) * 0.9
|
||||
)
|
||||
percentile_10 = distr.quantile_function(
|
||||
torch.ones((1, 1, 1)), torch.ones((1, 1)) * 0.1
|
||||
)
|
||||
|
||||
return samples.mean(), samples.std(), percentile_10, percentile_90
|
||||
|
||||
@@ -86,8 +84,8 @@ def test_independent_implicit_quantile() -> None:
|
||||
num_samples = NUM_SAMPLES
|
||||
|
||||
# # Normal distrib
|
||||
distr_mean = torch.Tensor([10.])
|
||||
distr_std = torch.Tensor([4.])
|
||||
distr_mean = torch.Tensor([10.0])
|
||||
distr_std = torch.Tensor([4.0])
|
||||
distr_pp10 = distr_mean - 1.282 * distr_std
|
||||
distr_pp90 = distr_mean + 1.282 * distr_std
|
||||
distr = Normal(loc=distr_mean, scale=distr_std)
|
||||
@@ -97,21 +95,29 @@ def test_independent_implicit_quantile() -> None:
|
||||
ImplicitQuantileOutput(output_domain="Real"),
|
||||
samples=samples,
|
||||
num_epochs=50,
|
||||
learning_rate=1e-2
|
||||
learning_rate=1e-2,
|
||||
)
|
||||
|
||||
torch.testing.assert_allclose(learned_mean, distr_mean.squeeze(), rtol=0.1, atol=0.1*10)
|
||||
torch.testing.assert_allclose(learned_std, distr_std.squeeze(), rtol=0.1, atol=.1*4)
|
||||
torch.testing.assert_allclose(learned_pp90, distr_pp90.squeeze(), rtol=0.1, atol=.1 * 4)
|
||||
torch.testing.assert_allclose(learned_pp10, distr_pp10.squeeze(), rtol=0.1, atol=.1 * 4)
|
||||
torch.testing.assert_allclose(
|
||||
learned_mean, distr_mean.squeeze(), rtol=0.1, atol=0.1 * 10
|
||||
)
|
||||
torch.testing.assert_allclose(
|
||||
learned_std, distr_std.squeeze(), rtol=0.1, atol=0.1 * 4
|
||||
)
|
||||
torch.testing.assert_allclose(
|
||||
learned_pp90, distr_pp90.squeeze(), rtol=0.1, atol=0.1 * 4
|
||||
)
|
||||
torch.testing.assert_allclose(
|
||||
learned_pp10, distr_pp10.squeeze(), rtol=0.1, atol=0.1 * 4
|
||||
)
|
||||
|
||||
# Uniform distrib
|
||||
a = torch.Tensor([0.])
|
||||
b = torch.Tensor([20.])
|
||||
distr_mean = 0.5*(a+b)
|
||||
distr_std = (1./12.*(b-a)**2)**0.5
|
||||
distr_pp10 = 0.1 * (a+b)
|
||||
distr_pp90 = 0.9 * (a+b)
|
||||
a = torch.Tensor([0.0])
|
||||
b = torch.Tensor([20.0])
|
||||
distr_mean = 0.5 * (a + b)
|
||||
distr_std = (1.0 / 12.0 * (b - a) ** 2) ** 0.5
|
||||
distr_pp10 = 0.1 * (a + b)
|
||||
distr_pp90 = 0.9 * (a + b)
|
||||
distr = Uniform(low=a, high=b)
|
||||
|
||||
samples = distr.sample((num_samples,))
|
||||
@@ -119,19 +125,25 @@ def test_independent_implicit_quantile() -> None:
|
||||
ImplicitQuantileOutput(output_domain="Positive"),
|
||||
samples=samples,
|
||||
num_epochs=50,
|
||||
learning_rate=1e-2
|
||||
learning_rate=1e-2,
|
||||
)
|
||||
|
||||
torch.testing.assert_allclose(learned_mean, distr_mean.squeeze(), atol=1., rtol=0.1)
|
||||
torch.testing.assert_allclose(
|
||||
learned_mean, distr_mean.squeeze(), atol=1.0, rtol=0.1
|
||||
)
|
||||
torch.testing.assert_allclose(learned_std, distr_std.squeeze(), atol=0.5, rtol=0.1)
|
||||
torch.testing.assert_allclose(learned_pp90, distr_pp90.squeeze(), rtol=0.1, atol=.1 * 18)
|
||||
torch.testing.assert_allclose(learned_pp10, distr_pp10.squeeze(), rtol=0.2, atol=.2 * 2)
|
||||
torch.testing.assert_allclose(
|
||||
learned_pp90, distr_pp90.squeeze(), rtol=0.1, atol=0.1 * 18
|
||||
)
|
||||
torch.testing.assert_allclose(
|
||||
learned_pp10, distr_pp10.squeeze(), rtol=0.2, atol=0.2 * 2
|
||||
)
|
||||
|
||||
# Bernoulli distrib
|
||||
distr_mean = torch.Tensor([0.2])
|
||||
distr_std = distr_mean * (1 - distr_mean)
|
||||
distr_pp10 = torch.Tensor([0.])
|
||||
distr_pp90 = torch.Tensor([1.])
|
||||
distr_pp10 = torch.Tensor([0.0])
|
||||
distr_pp90 = torch.Tensor([1.0])
|
||||
distr = Bernoulli(probs=distr_mean)
|
||||
|
||||
samples = distr.sample((num_samples,))
|
||||
@@ -139,13 +151,19 @@ def test_independent_implicit_quantile() -> None:
|
||||
ImplicitQuantileOutput(output_domain="Positive"),
|
||||
samples=samples,
|
||||
num_epochs=50,
|
||||
learning_rate=1e-2
|
||||
learning_rate=1e-2,
|
||||
)
|
||||
|
||||
torch.testing.assert_allclose(learned_mean, distr_mean.squeeze(), atol=1., rtol=0.1)
|
||||
torch.testing.assert_allclose(
|
||||
learned_mean, distr_mean.squeeze(), atol=1.0, rtol=0.1
|
||||
)
|
||||
torch.testing.assert_allclose(learned_std, distr_std.squeeze(), atol=0.5, rtol=0.1)
|
||||
torch.testing.assert_allclose(learned_pp90, distr_pp90.squeeze(), rtol=0.1, atol=.1 * 18)
|
||||
torch.testing.assert_allclose(learned_pp10, distr_pp10.squeeze(), rtol=0.1, atol=.1 * 2)
|
||||
torch.testing.assert_allclose(
|
||||
learned_pp90, distr_pp90.squeeze(), rtol=0.1, atol=0.1 * 18
|
||||
)
|
||||
torch.testing.assert_allclose(
|
||||
learned_pp10, distr_pp10.squeeze(), rtol=0.1, atol=0.1 * 2
|
||||
)
|
||||
|
||||
|
||||
def test_training_with_implicit_quantile_output():
|
||||
@@ -156,16 +174,16 @@ def test_training_with_implicit_quantile_output():
|
||||
distr_output=ImplicitQuantileOutput(output_domain="Real"),
|
||||
freq=metadata.freq,
|
||||
prediction_length=metadata.prediction_length,
|
||||
trainer=Trainer(device="cpu",
|
||||
epochs=5,
|
||||
learning_rate=1e-3,
|
||||
num_batches_per_epoch=3,
|
||||
batch_size=256,
|
||||
num_workers=1,
|
||||
),
|
||||
trainer=Trainer(
|
||||
device="cpu",
|
||||
epochs=5,
|
||||
learning_rate=1e-3,
|
||||
num_batches_per_epoch=3,
|
||||
batch_size=256,
|
||||
),
|
||||
input_size=48,
|
||||
)
|
||||
deepar_predictor = deepar_estimator.train(dataset.train)
|
||||
deepar_predictor = deepar_estimator.train(dataset.train, num_workers=1)
|
||||
forecast_it, ts_it = make_evaluation_predictions(
|
||||
dataset=dataset.test, # test dataset
|
||||
predictor=deepar_predictor, # predictor
|
||||
@@ -174,13 +192,14 @@ def test_training_with_implicit_quantile_output():
|
||||
forecasts = list(forecast_it)
|
||||
tss = list(ts_it)
|
||||
evaluator = Evaluator(num_workers=0)
|
||||
agg_metrics, item_metrics = evaluator(iter(tss), iter(forecasts), num_series=len(dataset.test))
|
||||
agg_metrics, item_metrics = evaluator(
|
||||
iter(tss), iter(forecasts), num_series=len(dataset.test)
|
||||
)
|
||||
|
||||
assert agg_metrics["MSE"] > 0
|
||||
|
||||
|
||||
def test_instanciation_of_args_proj():
|
||||
|
||||
class MockedImplicitQuantileOutput(ImplicitQuantileOutput):
|
||||
method_calls = 0
|
||||
|
||||
@@ -198,17 +217,17 @@ def test_instanciation_of_args_proj():
|
||||
distr_output=distr_output,
|
||||
freq=metadata.freq,
|
||||
prediction_length=metadata.prediction_length,
|
||||
trainer=Trainer(device="cpu",
|
||||
epochs=1,
|
||||
learning_rate=1e-3,
|
||||
num_batches_per_epoch=1,
|
||||
batch_size=256,
|
||||
num_workers=1,
|
||||
),
|
||||
trainer=Trainer(
|
||||
device="cpu",
|
||||
epochs=3,
|
||||
learning_rate=1e-3,
|
||||
num_batches_per_epoch=1,
|
||||
batch_size=256,
|
||||
),
|
||||
input_size=48,
|
||||
)
|
||||
assert distr_output.method_calls == 1
|
||||
deepar_predictor = deepar_estimator.train(dataset.train)
|
||||
deepar_predictor = deepar_estimator.train(dataset.train, num_workers=1)
|
||||
|
||||
# Method should be called when the MockedImplicitQuantileOutput is instanciated,
|
||||
# and one more time because in_features is different from 1
|
||||
@@ -222,7 +241,9 @@ def test_instanciation_of_args_proj():
|
||||
forecasts = list(forecast_it)
|
||||
tss = list(ts_it)
|
||||
evaluator = Evaluator(num_workers=0)
|
||||
agg_metrics, item_metrics = evaluator(iter(tss), iter(forecasts), num_series=len(dataset.test))
|
||||
agg_metrics, item_metrics = evaluator(
|
||||
iter(tss), iter(forecasts), num_series=len(dataset.test)
|
||||
)
|
||||
assert distr_output.method_calls == 2
|
||||
|
||||
# Test that the implicit output module is proper reset
|
||||
@@ -230,15 +251,17 @@ def test_instanciation_of_args_proj():
|
||||
distr_output=MockedImplicitQuantileOutput(output_domain="Real"),
|
||||
freq=metadata.freq,
|
||||
prediction_length=metadata.prediction_length,
|
||||
trainer=Trainer(device="cpu",
|
||||
epochs=1,
|
||||
learning_rate=1e-3,
|
||||
num_batches_per_epoch=1,
|
||||
batch_size=256,
|
||||
num_workers=1,
|
||||
),
|
||||
trainer=Trainer(
|
||||
device="cpu",
|
||||
epochs=3,
|
||||
learning_rate=1e-3,
|
||||
num_batches_per_epoch=1,
|
||||
batch_size=256,
|
||||
),
|
||||
input_size=48,
|
||||
)
|
||||
assert distr_output.method_calls == 3
|
||||
new_estimator.train(dataset.train)
|
||||
assert distr_output.method_calls == 3 # Since in_feature is the same as before, there should be no additional call
|
||||
new_estimator.train(dataset.train, num_workers=1)
|
||||
assert (
|
||||
distr_output.method_calls == 3
|
||||
) # Since in_feature is the same as before, there should be no additional call
|
||||
|
||||
Reference in New Issue
Block a user