initial torchscale models

This commit is contained in:
Kashif Rasul committed 2022-11-28 19:51:58 +01:00
1 parent efc4d0d678
commit f6e7d77fe4
6 files changed
+1679

No files matched your search

+2
View File
@@ -11,3 +11,5 @@ einops
opt_einsum
pykeops
scipy
apex
torchscale
+10
View File
@@ -0,0 +1,10 @@
# +
from .estimator import TorchscaleEstimator
from .lightning_module import TorchscaleightningModule
from .module import TorchscaleModel
__all__ = [
"TorchscaleModel",
"TorchscaleightningModule",
"TorchscaleEstimator",
]
+295
View File
@@ -0,0 +1,295 @@
# +
from typing import Any, Dict, Iterable, List, Optional
import torch
from torch.utils.data import DataLoader
import numpy as np
from gluonts.core.component import validated
from gluonts.dataset.common import Dataset
from gluonts.dataset.field_names import FieldName
from gluonts.itertools import Cyclic, IterableSlice, PseudoShuffled
from gluonts.time_feature import TimeFeature, time_features_from_frequency_str
from gluonts.torch.model.estimator import PyTorchLightningEstimator
from gluonts.torch.model.predictor import PyTorchPredictor
from gluonts.torch.distributions import DistributionOutput, StudentTOutput
from gluonts.torch.modules.loss import DistributionLoss, NegativeLogLikelihood
from gluonts.torch.util import IterableDataset
from gluonts.transform import (
AddAgeFeature,
AddObservedValuesIndicator,
AddTimeFeatures,
AsNumpyArray,
Chain,
ExpectedNumInstanceSampler,
InstanceSplitter,
RemoveFields,
SelectFields,
SetField,
TestSplitSampler,
Transformation,
ValidationSplitSampler,
VstackFeatures,
)
from torchscale.architecture.config import EncoderDecoderConfig
from lightning_module import TorchscaleLightningModule
from module import TorchscaleModel
# +
PREDICTION_INPUT_NAMES = [
"feat_static_cat",
"feat_static_real",
"past_time_feat",
"past_target",
"past_observed_values",
"future_time_feat",
]
TRAINING_INPUT_NAMES = PREDICTION_INPUT_NAMES + [
"future_target",
"future_observed_values",
]
class TorchscaleEstimator(PyTorchLightningEstimator):
@validated()
def __init__(
self,
freq: str,
prediction_length: int,
# Torchscale arguments
enc_dec_config: EncoderDecoderConfig,
input_size: int = 1,
context_length: Optional[int] = None,
num_feat_dynamic_real: int = 0,
num_feat_static_cat: int = 0,
num_feat_static_real: int = 0,
cardinality: Optional[List[int]] = None,
embedding_dimension: Optional[List[int]] = None,
distr_output: DistributionOutput = StudentTOutput(),
loss: DistributionLoss = NegativeLogLikelihood(),
scaling: bool = True,
lags_seq: Optional[List[int]] = None,
time_features: Optional[List[TimeFeature]] = None,
num_parallel_samples: int = 100,
batch_size: int = 32,
num_batches_per_epoch: int = 50,
trainer_kwargs: Optional[Dict[str, Any]] = dict(),
) -> None:
trainer_kwargs = {
"max_epochs": 100,
**trainer_kwargs,
}
super().__init__(trainer_kwargs=trainer_kwargs)
self.freq = freq
self.context_length = (
context_length if context_length is not None else prediction_length
)
self.prediction_length = prediction_length
self.distr_output = distr_output
self.loss = loss
self.enc_dec_config = enc_dec_config
self.input_size = input_size
self.num_feat_dynamic_real = num_feat_dynamic_real
self.num_feat_static_cat = num_feat_static_cat
self.num_feat_static_real = num_feat_static_real
self.cardinality = (
cardinality if cardinality and num_feat_static_cat > 0 else [1]
)
self.embedding_dimension = embedding_dimension
self.scaling = scaling
self.lags_seq = lags_seq
self.time_features = (
time_features
if time_features is not None
else time_features_from_frequency_str(self.freq)
)
self.num_parallel_samples = num_parallel_samples
self.batch_size = batch_size
self.num_batches_per_epoch = num_batches_per_epoch
self.train_sampler = ExpectedNumInstanceSampler(
num_instances=1.0, min_future=prediction_length
)
self.validation_sampler = ValidationSplitSampler(min_future=prediction_length)
def create_transformation(self) -> Transformation:
remove_field_names = []
if self.num_feat_static_real == 0:
remove_field_names.append(FieldName.FEAT_STATIC_REAL)
if self.num_feat_dynamic_real == 0:
remove_field_names.append(FieldName.FEAT_DYNAMIC_REAL)
return Chain(
[RemoveFields(field_names=remove_field_names)]
+ (
[SetField(output_field=FieldName.FEAT_STATIC_CAT, value=[0])]
if not self.num_feat_static_cat > 0
else []
)
+ (
[SetField(output_field=FieldName.FEAT_STATIC_REAL, value=[0.0])]
if not self.num_feat_static_real > 0
else []
)
+ [
AsNumpyArray(
field=FieldName.FEAT_STATIC_CAT,
expected_ndim=1,
dtype=np.long,
),
AsNumpyArray(
field=FieldName.FEAT_STATIC_REAL,
expected_ndim=1,
),
AsNumpyArray(
field=FieldName.TARGET,
# in the following line, we add 1 for the time dimension
expected_ndim=1 + len(self.distr_output.event_shape),
),
AddObservedValuesIndicator(
target_field=FieldName.TARGET,
output_field=FieldName.OBSERVED_VALUES,
),
AddTimeFeatures(
start_field=FieldName.START,
target_field=FieldName.TARGET,
output_field=FieldName.FEAT_TIME,
time_features=self.time_features,
pred_length=self.prediction_length,
),
AddAgeFeature(
target_field=FieldName.TARGET,
output_field=FieldName.FEAT_AGE,
pred_length=self.prediction_length,
log_scale=True,
),
VstackFeatures(
output_field=FieldName.FEAT_TIME,
input_fields=[FieldName.FEAT_TIME, FieldName.FEAT_AGE]
+ (
[FieldName.FEAT_DYNAMIC_REAL]
if self.num_feat_dynamic_real > 0
else []
),
),
]
)
def _create_instance_splitter(self, module: TorchscaleLightningModule, mode: str):
assert mode in ["training", "validation", "test"]
instance_sampler = {
"training": self.train_sampler,
"validation": self.validation_sampler,
"test": TestSplitSampler(),
}[mode]
return InstanceSplitter(
target_field=FieldName.TARGET,
is_pad_field=FieldName.IS_PAD,
start_field=FieldName.START,
forecast_start_field=FieldName.FORECAST_START,
instance_sampler=instance_sampler,
past_length=module.model._past_length,
future_length=self.prediction_length,
time_series_fields=[
FieldName.FEAT_TIME,
FieldName.OBSERVED_VALUES,
],
dummy_value=self.distr_output.value_in_support,
)
def create_training_data_loader(
self,
data: Dataset,
module: TorchscaleLightningModule,
shuffle_buffer_length: Optional[int] = None,
**kwargs,
) -> Iterable:
transformation = self._create_instance_splitter(
module, "training"
) + SelectFields(TRAINING_INPUT_NAMES)
training_instances = transformation.apply(
Cyclic(data)
if shuffle_buffer_length is None
else PseudoShuffled(
Cyclic(data), shuffle_buffer_length=shuffle_buffer_length
)
)
return IterableSlice(
iter(
DataLoader(
IterableDataset(training_instances),
batch_size=self.batch_size,
**kwargs,
)
),
self.num_batches_per_epoch,
)
def create_validation_data_loader(
self,
data: Dataset,
module: TorchscaleLightningModule,
**kwargs,
) -> Iterable:
transformation = self._create_instance_splitter(
module, "validation"
) + SelectFields(TRAINING_INPUT_NAMES)
validation_instances = transformation.apply(data)
return DataLoader(
IterableDataset(validation_instances),
batch_size=self.batch_size,
**kwargs,
)
def create_predictor(
self,
transformation: Transformation,
module: TorchscaleLightningModule,
) -> PyTorchPredictor:
prediction_splitter = self._create_instance_splitter(module, "test")
return PyTorchPredictor(
input_transform=transformation + prediction_splitter,
input_names=PREDICTION_INPUT_NAMES,
prediction_net=module.model,
batch_size=self.batch_size,
prediction_length=self.prediction_length,
device=torch.device("cuda" if torch.cuda.is_available() else "cpu"),
)
def create_lightning_module(self) -> TorchscaleLightningModule:
model = TorchscaleModel(
freq=self.freq,
context_length=self.context_length,
prediction_length=self.prediction_length,
num_feat_dynamic_real=1
+ self.num_feat_dynamic_real
+ len(self.time_features),
num_feat_static_real=max(1, self.num_feat_static_real),
num_feat_static_cat=max(1, self.num_feat_static_cat),
cardinality=self.cardinality,
embedding_dimension=self.embedding_dimension,
# torchscale configs
enc_dec_config=self.enc_dec_config,
# univariate input
input_size=self.input_size,
distr_output=self.distr_output,
lags_seq=self.lags_seq,
scaling=self.scaling,
num_parallel_samples=self.num_parallel_samples,
)
return TorchscaleLightningModule(model=model, loss=self.loss)
+81
View File
@@ -0,0 +1,81 @@
import pytorch_lightning as pl
import torch
from gluonts.torch.modules.loss import DistributionLoss, NegativeLogLikelihood
from gluonts.torch.util import weighted_average
from module import TorchscaleModel
class TorchscaleLightningModule(pl.LightningModule):
def __init__(
self,
model: TorchscaleModel,
loss: DistributionLoss = NegativeLogLikelihood(),
lr: float = 5e-3,
weight_decay: float = 1e-6,
) -> None:
super().__init__()
self.save_hyperparameters()
self.model = model
self.loss = loss
self.lr = lr
self.weight_decay = weight_decay
def training_step(self, batch, batch_idx: int):
"""Execute training step"""
train_loss = self(batch)
self.log(
"train_loss",
train_loss,
on_epoch=True,
on_step=False,
prog_bar=True,
)
return train_loss
def validation_step(self, batch, batch_idx: int):
"""Execute validation step"""
with torch.no_grad():
val_loss = self(batch)
self.log("val_loss", val_loss, on_epoch=True, on_step=False, prog_bar=True)
return val_loss
def configure_optimizers(self):
"""Returns the optimizer to use"""
return torch.optim.Adam(
self.model.parameters(),
lr=self.lr,
weight_decay=self.weight_decay,
)
def forward(self, batch):
feat_static_cat = batch["feat_static_cat"]
feat_static_real = batch["feat_static_real"]
past_time_feat = batch["past_time_feat"]
past_target = batch["past_target"]
future_time_feat = batch["future_time_feat"]
future_target = batch["future_target"]
past_observed_values = batch["past_observed_values"]
future_observed_values = batch["future_observed_values"]
transformer_inputs, scale, _ = self.model.create_network_inputs(
feat_static_cat,
feat_static_real,
past_time_feat,
past_target,
past_observed_values,
future_time_feat,
future_target,
)
params = self.model.output_params(transformer_inputs)
distr = self.model.output_distribution(params, scale)
loss_values = self.loss(distr, future_target)
if len(self.model.target_shape) == 0:
loss_weights = future_observed_values
else:
loss_weights, _ = future_observed_values.min(dim=-1, keepdim=False)
return weighted_average(loss_values, weights=loss_weights)
+623
View File
@@ -0,0 +1,623 @@
from typing import List, Optional, Dict, Any
import math
import torch
import torch.nn as nn
from gluonts.core.component import validated
from gluonts.time_feature import get_lags_for_frequency
from gluonts.torch.distributions import DistributionOutput, StudentTOutput
from gluonts.torch.modules.feature import FeatureEmbedder
from gluonts.torch.modules.scaler import MeanScaler, NOPScaler
from apex.normalization import FusedLayerNorm as LayerNorm
from torchscale.architecture.config import EncoderDecoderConfig
from torchscale.component.relative_position_bias import RelativePositionBias
from torchscale.architecture.encoder import EncoderLayer
from torchscale.architecture.decoder import DecoderLayer
from torchscale.component.multiway_network import MultiwayWrapper
from torchscale.architecture.utils import init_bert_params
class Encoder(nn.Module):
def __init__(self, args, is_moe_layer=False, is_encoder_decoder=True):
super().__init__()
self.dropout_module = torch.nn.Dropout(args.dropout, inplace=True)
embed_dim = args.encoder_embed_dim
self.layers = nn.ModuleList([])
moe_freq = args.moe_freq
for i in range(args.encoder_layers):
is_moe_layer = moe_freq != 0 and (i + 1) % moe_freq == 0
self.layers.append(
self.build_encoder_layer(
args,
depth=i,
is_moe_layer=is_moe_layer,
is_encoder_decoder=is_encoder_decoder,
)
)
self.num_layers = len(self.layers)
if args.encoder_normalize_before:
self.layer_norm = MultiwayWrapper(args, LayerNorm(embed_dim))
else:
self.layer_norm = None
if args.rel_pos_buckets > 0 and args.max_rel_pos > 0:
self.relative_position = RelativePositionBias(
num_buckets=args.rel_pos_buckets,
max_distance=args.max_rel_pos,
n_heads=args.encoder_attention_heads,
)
else:
self.relative_position = None
if args.bert_init:
self.apply(init_bert_params)
if args.deepnorm:
if is_encoder_decoder:
init_scale = (
math.pow(
math.pow(args.encoder_layers, 4) * args.decoder_layers, 0.0625
)
/ 1.15
)
else:
init_scale = math.pow(8.0 * args.encoder_layers, 0.25)
for name, p in self.named_parameters():
if (
"fc1" in name
or "fc2" in name
or "out_proj" in name
or "v_proj" in name
):
p.data.div_(init_scale)
if args.subln:
if is_encoder_decoder:
init_scale = math.sqrt(
math.log(3 * args.decoder_layers)
* math.log(2 * args.encoder_layers)
/ 3
)
else:
init_scale = math.sqrt(math.log(args.encoder_layers * 2))
for name, p in self.named_parameters():
if (
"fc1" in name
or "fc2" in name
or "out_proj" in name
or "v_proj" in name
):
p.data.mul_(init_scale)
def build_encoder_layer(
self, args, depth, is_moe_layer=False, is_encoder_decoder=False
):
layer = EncoderLayer(
args,
depth,
is_moe_layer=is_moe_layer,
is_encoder_decoder=is_encoder_decoder,
)
return layer
def forward(self, enc_input, encoder_padding_mask=None):
x = enc_input.transpose(0, 1) # (B, T, C) -> (T, B, C)
rel_pos_bias = None
if self.relative_position is not None:
rel_pos_bias = self.relative_position(
batch_size=x.size(1), qlen=x.size(0), klen=x.size(0)
)
for layer in self.layers:
x, _ = layer(
x, encoder_padding_mask=encoder_padding_mask, rel_pos=rel_pos_bias
)
if self.layer_norm is not None:
x = self.layer_norm(x)
return x # (T, B, C)
class Decoder(nn.Module):
def __init__(self, args, is_encoder_decoder=True):
super().__init__()
embed_dim = args.decoder_embed_dim
self.dropout_module = torch.nn.Dropout(args.dropout, inplace=True)
if args.layernorm_embedding:
self.layernorm_embedding = LayerNorm(embed_dim)
else:
self.layernorm_embedding = None
self.layers = nn.ModuleList([])
moe_freq = args.moe_freq
for i in range(args.decoder_layers):
is_moe_layer = moe_freq != 0 and (i + 1) % moe_freq == 0
self.layers.append(
self.build_decoder_layer(
args,
depth=i,
is_moe_layer=is_moe_layer,
is_encoder_decoder=is_encoder_decoder,
)
)
self.num_layers = len(self.layers)
if args.decoder_normalize_before:
self.layer_norm = LayerNorm(embed_dim)
else:
self.layer_norm = None
self.self_attn_relative_position = None
self.cross_attn_relative_position = None
if args.rel_pos_buckets > 0 and args.max_rel_pos > 0:
self.self_attn_relative_position = RelativePositionBias(
num_buckets=args.rel_pos_buckets,
max_distance=args.max_rel_pos,
n_heads=args.decoder_attention_heads,
)
if is_encoder_decoder:
self.cross_attn_relative_position = RelativePositionBias(
num_buckets=args.rel_pos_buckets,
max_distance=args.max_rel_pos,
n_heads=args.decoder_attention_heads,
)
if args.bert_init:
self.apply(init_bert_params)
if args.deepnorm:
if is_encoder_decoder:
init_scale = math.pow(12.0 * args.decoder_layers, 0.25)
else:
init_scale = math.pow(8.0 * args.decoder_layers, 0.25)
for name, p in self.named_parameters():
if (
"fc1" in name
or "fc2" in name
or "out_proj" in name
or "v_proj" in name
):
p.data.div_(init_scale)
if args.subln:
if is_encoder_decoder:
init_scale = math.sqrt(math.log(args.decoder_layers * 3))
else:
init_scale = math.sqrt(math.log(args.decoder_layers * 2))
for name, p in self.named_parameters():
if "encoder_attn" in name:
continue
if (
"fc1" in name
or "fc2" in name
or "out_proj" in name
or "v_proj" in name
):
p.data.mul_(init_scale)
def build_decoder_layer(
self, args, depth, is_moe_layer=False, is_encoder_decoder=False
):
layer = DecoderLayer(
args,
depth,
is_moe_layer=is_moe_layer,
is_encoder_decoder=is_encoder_decoder,
)
return layer
def forward(self, dec_input, encoder_out, incremental_state=None):
x = dec_input.transpose(0, 1) # (B, T, C) -> (T, B, C)
# relative position
self_attn_rel_pos_bias = None
slen = dec_input.size(1)
if self.self_attn_relative_position is not None:
self_attn_rel_pos_bias = self.self_attn_relative_position(
batch_size=x.size(1), qlen=slen, klen=slen
)
if incremental_state is not None:
self_attn_rel_pos_bias = self_attn_rel_pos_bias[:, -1:, :]
cross_attn_rel_pos_bias = None
if self.cross_attn_relative_position is not None:
cross_attn_rel_pos_bias = self.cross_attn_relative_position(
batch_size=x.size(1),
qlen=slen,
klen=encoder_out["encoder_out"].size(0),
)
if incremental_state is not None:
cross_attn_rel_pos_bias = cross_attn_rel_pos_bias[:, -1:, :]
# decoder layers
for idx, layer in enumerate(self.layers):
if incremental_state is None:
self_attn_mask = torch.triu(
torch.zeros([x.size(0), x.size(0)])
.float()
.fill_(float("-inf"))
.type_as(x),
1,
)
else:
self_attn_mask = None
if idx not in incremental_state:
incremental_state[idx] = {}
x, _, _, _ = layer(
x,
encoder_out,
None,
incremental_state[idx] if incremental_state is not None else None,
self_attn_mask=self_attn_mask,
self_attn_padding_mask=None,
self_attn_rel_pos=self_attn_rel_pos_bias,
cross_attn_rel_pos=cross_attn_rel_pos_bias,
)
if self.layer_norm is not None:
x = self.layer_norm(x)
return x.transpose(0, 1)
class TorchscaleModel(nn.Module):
@validated()
def __init__(
self,
freq: str,
context_length: int,
prediction_length: int,
num_feat_dynamic_real: int,
num_feat_static_real: int,
num_feat_static_cat: int,
cardinality: List[int],
# torchscale config
enc_dec_config: EncoderDecoderConfig,
input_size: int = 1,
embedding_dimension: Optional[List[int]] = None,
distr_output: DistributionOutput = StudentTOutput(),
lags_seq: Optional[List[int]] = None,
scaling: bool = True,
num_parallel_samples: int = 1,
) -> None:
super().__init__()
self.input_size = input_size
self.target_shape = distr_output.event_shape
self.num_feat_dynamic_real = num_feat_dynamic_real
self.num_feat_static_cat = num_feat_static_cat
self.num_feat_static_real = num_feat_static_real
self.embedding_dimension = (
embedding_dimension
if embedding_dimension is not None or cardinality is None
else [min(50, (cat + 1) // 2) for cat in cardinality]
)
self.lags_seq = lags_seq or get_lags_for_frequency(freq_str=freq)
self.num_parallel_samples = num_parallel_samples
self.history_length = context_length + max(self.lags_seq)
self.embedder = FeatureEmbedder(
cardinalities=cardinality,
embedding_dims=self.embedding_dimension,
)
if scaling:
self.scaler = MeanScaler(dim=1, keepdim=True)
else:
self.scaler = NOPScaler(dim=1, keepdim=True)
# total feature size
d_model = self.input_size * len(self.lags_seq) + self._number_of_features
self.context_length = context_length
self.prediction_length = prediction_length
self.distr_output = distr_output
self.param_proj = distr_output.get_args_proj(d_model)
enc_dec_config.encoder_embed_dim = d_model
enc_dec_config.decoder_embed_dim = d_model
self.encoder = Encoder(enc_dec_config)
self.decoder = Decoder(enc_dec_config)
# attention_args["dropout"] = dropout
# attention_args["causal"] = False
# attention_args["seq_len"] = self.context_length
# attention_args["num_rules"] = nhead
# attention_args["attention_query_mask"] = torch.rand((context_length, 1)) < 0.5
# xformer_config = [
# # A list of the encoder blocks which constitute the Transformer.
# # Note that a sequence of different encoder blocks can be used
# {
# "reversible": reversible, # Optionally make these layers reversible, to save memory
# "block_type": "encoder",
# "num_layers": num_encoder_layers, # Optional, this means that this config will repeat N times
# "dim_model": d_model,
# "residual_norm_style": residual_norm_style, # Optional, pre/post
# "position_encoding_config": {
# "name": "sine",
# "dim_model": d_model,
# },
# "multi_head_config": {
# "use_rotary_embeddings": use_rotary_embeddings,
# "num_heads": nhead,
# "residual_dropout": dropout,
# "attention": attention_args,
# },
# "feedforward_config": {
# "name": "MLP",
# "dropout": dropout,
# "activation": activation,
# "hidden_layer_multiplier": hidden_layer_multiplier,
# "dim_model": d_model,
# },
# },
# ]
# config = xFormerConfig(xformer_config)
# # xformer encoder
# self.encoder = xFormer.from_config(config)
# # causal vanilla transformer decoder
# decoder_layer = nn.TransformerDecoderLayer(
# d_model,
# nhead,
# dim_feedforward=d_model * hidden_layer_multiplier,
# dropout=dropout,
# activation=activation,
# layer_norm_eps=1e-5,
# batch_first=True,
# norm_first=False,
# )
# decoder_norm = nn.LayerNorm(d_model, eps=1e-5)
# self.decoder = nn.TransformerDecoder(
# decoder_layer, num_decoder_layers, decoder_norm
# )
# causal decoder tgt mask for training
self.register_buffer(
"tgt_mask",
nn.Transformer.generate_square_subsequent_mask(prediction_length),
)
@property
def _number_of_features(self) -> int:
return (
sum(self.embedding_dimension)
+ self.num_feat_dynamic_real
+ self.num_feat_static_real
+ self.input_size # the log(scale)
)
@property
def _past_length(self) -> int:
return self.context_length + max(self.lags_seq)
def get_lagged_subsequences(
self, sequence: torch.Tensor, subsequences_length: int, shift: int = 0
) -> torch.Tensor:
"""
Returns lagged subsequences of a given sequence.
Parameters
----------
sequence : Tensor
the sequence from which lagged subsequences should be extracted.
Shape: (N, T, C).
subsequences_length : int
length of the subsequences to be extracted.
shift: int
shift the lags by this amount back.
Returns
--------
lagged : Tensor
a tensor of shape (N, S, C, I), where S = subsequences_length and
I = len(indices), containing lagged subsequences. Specifically,
lagged[i, j, :, k] = sequence[i, -indices[k]-S+j, :].
"""
sequence_length = sequence.shape[1]
indices = [l - shift for l in self.lags_seq]
assert max(indices) + subsequences_length <= sequence_length, (
f"lags cannot go further than history length, found lag {max(indices)} "
f"while history length is only {sequence_length}"
)
lagged_values = []
for lag_index in indices:
begin_index = -lag_index - subsequences_length
end_index = -lag_index if lag_index > 0 else None
lagged_values.append(sequence[:, begin_index:end_index, ...])
return torch.stack(lagged_values, dim=-1)
def create_network_inputs(
self,
feat_static_cat: torch.Tensor,
feat_static_real: torch.Tensor,
past_time_feat: torch.Tensor,
past_target: torch.Tensor,
past_observed_values: torch.Tensor,
future_time_feat: Optional[torch.Tensor] = None,
future_target: Optional[torch.Tensor] = None,
):
# time feature
time_feat = (
past_time_feat[:, self._past_length - self.context_length :, ...]
if future_time_feat is None or future_target is None
else torch.cat(
(
past_time_feat[:, self._past_length - self.context_length :, ...],
future_time_feat,
),
dim=1,
)
)
# target
context = past_target[:, -self.context_length :]
observed_context = past_observed_values[:, -self.context_length :]
# weights = torch.linspace(0.0001, 1, steps=observed_context.size(-1), device=observed_context.device)
_, scale = self.scaler(context, observed_context)
inputs = (
torch.cat((past_target, future_target), dim=1) / scale
if future_target is not None
else past_target / scale
)
inputs_length = (
self._past_length + self.prediction_length
if future_target is not None
else self._past_length
)
assert inputs.shape[1] == inputs_length
subsequences_length = (
self.context_length
if future_time_feat is None or future_target is None
else self.context_length + self.prediction_length
)
# embeddings
embedded_cat = self.embedder(feat_static_cat)
log_scale = scale.log() if self.input_size == 1 else scale.squeeze(1).log()
static_feat = torch.cat(
(embedded_cat, feat_static_real, log_scale),
dim=1,
)
expanded_static_feat = static_feat.unsqueeze(1).expand(
-1, time_feat.shape[1], -1
)
features = torch.cat((expanded_static_feat, time_feat), dim=-1)
# self._check_shapes(prior_input, inputs, features)
# sequence = torch.cat((prior_input, inputs), dim=1)
lagged_sequence = self.get_lagged_subsequences(
sequence=inputs,
subsequences_length=subsequences_length,
)
lags_shape = lagged_sequence.shape
reshaped_lagged_sequence = lagged_sequence.reshape(
lags_shape[0], lags_shape[1], -1
)
if features is None:
transformer_inputs = reshaped_lagged_sequence
else:
transformer_inputs = torch.cat((reshaped_lagged_sequence, features), dim=-1)
return transformer_inputs, scale, static_feat
def output_params(self, transformer_inputs):
enc_input = transformer_inputs[:, : self.context_length, ...]
dec_input = transformer_inputs[:, self.context_length :, ...]
enc_out = self.encoder(enc_input)
dec_output = self.decoder(dec_input, enc_out)
return self.param_proj(dec_output)
@torch.jit.ignore
def output_distribution(
self, params, scale=None, trailing_n=None
) -> torch.distributions.Distribution:
sliced_params = params
if trailing_n is not None:
sliced_params = [p[:, -trailing_n:] for p in params]
return self.distr_output.distribution(sliced_params, scale=scale)
# for prediction
def forward(
self,
feat_static_cat: torch.Tensor,
feat_static_real: torch.Tensor,
past_time_feat: torch.Tensor,
past_target: torch.Tensor,
past_observed_values: torch.Tensor,
future_time_feat: torch.Tensor,
num_parallel_samples: Optional[int] = None,
) -> torch.Tensor:
if num_parallel_samples is None:
num_parallel_samples = self.num_parallel_samples
encoder_inputs, scale, static_feat = self.create_network_inputs(
feat_static_cat,
feat_static_real,
past_time_feat,
past_target,
past_observed_values,
future_time_feat,
)
enc_out = self.encoder(src=encoder_inputs)
params = self.param_proj(enc_out.transpose(0, 1)) # (B, T, D)
distr = self.output_distribution(params, trailing_n=1)
repeated_scale = scale.repeat_interleave(
repeats=self.num_parallel_samples, dim=0
)
repeated_static_feat = static_feat.repeat_interleave(
repeats=self.num_parallel_samples, dim=0
).unsqueeze(dim=1)
repeated_past_target = (
past_target.repeat_interleave(repeats=self.num_parallel_samples, dim=0)
/ repeated_scale
)
repeated_time_feat = future_time_feat.repeat_interleave(
repeats=self.num_parallel_samples, dim=0
)
repeated_enc_out = enc_out.repeat_interleave(
repeats=self.num_parallel_samples, dim=0
)
future_samples = []
for k in range(self.prediction_length):
next_features = torch.cat(
(repeated_static_feat, repeated_time_feat[:, k : k + 1]),
dim=-1,
)
lagged_sequence = self.get_lagged_subsequences(
sequence=repeated_past_target,
subsequences_length=1,
shift=1,
)
lags_shape = lagged_sequence.shape
reshaped_lagged_sequence = lagged_sequence.reshape(
lags_shape[0], lags_shape[1], -1
)
decoder_input = torch.cat((reshaped_lagged_sequence, next_features), dim=-1)
output = self.decoder(decoder_input, repeated_enc_out)
params = self.param_proj(output)
distr = self.output_distribution(params)
next_sample = distr.sample()
repeated_past_target = torch.cat((repeated_past_target, next_sample), dim=1)
future_samples.append(next_sample)
unscaled_future_samples = torch.cat(future_samples, dim=1) * repeated_scale
return unscaled_future_samples.reshape(
(-1, self.num_parallel_samples, self.prediction_length) + self.target_shape,
)
+668
View File
@@ -0,0 +1,668 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": 1,
"id": "7c64affd",
"metadata": {},
"outputs": [],
"source": [
"from itertools import islice"
]
},
{
"cell_type": "code",
"execution_count": 2,
"id": "8aa55868",
"metadata": {},
"outputs": [],
"source": [
"%matplotlib inline\n",
"from matplotlib import pyplot as plt\n",
"import matplotlib.dates as mdates\n",
"\n",
"import pandas as pd\n",
"from sklearn.manifold import TSNE"
]
},
{
"cell_type": "code",
"execution_count": 3,
"id": "6a730716",
"metadata": {},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"2022-11-28 19:48:37.174077: I tensorflow/core/platform/cpu_feature_guard.cc:193] This TensorFlow binary is optimized with oneAPI Deep Neural Network Library (oneDNN) to use the following CPU instructions in performance-critical operations: SSE3 SSE4.1 SSE4.2 AVX AVX2 FMA\n",
"To enable them in other operations, rebuild TensorFlow with the appropriate compiler flags.\n"
]
}
],
"source": [
"from gluonts.dataset.repository.datasets import get_dataset\n",
"from gluonts.dataset.common import ListDataset\n",
"from gluonts.evaluation import make_evaluation_predictions, Evaluator\n",
"from pytorch_lightning.loggers import CSVLogger\n",
"from datasets import load_dataset\n",
"from torchscale.architecture.config import EncoderDecoderConfig\n",
"\n",
"from estimator import TorchscaleEstimator"
]
},
{
"cell_type": "code",
"execution_count": 4,
"id": "fc889c9f",
"metadata": {},
"outputs": [],
"source": [
"dataset = get_dataset(\"electricity\")"
]
},
{
"cell_type": "code",
"execution_count": 5,
"id": "666105d9",
"metadata": {},
"outputs": [],
"source": [
"enc_dec_config = EncoderDecoderConfig(\n",
" encoder_attention_heads=2,\n",
" decoder_attention_heads=2,\n",
" encoder_layers=4,\n",
" decoder_layers=4,\n",
" encoder_ffn_embed_dim=32,\n",
" decoder_ffn_embed_dim=32,\n",
")"
]
},
{
"cell_type": "code",
"execution_count": 6,
"id": "1717d0d2",
"metadata": {},
"outputs": [],
"source": [
"estimator = TorchscaleEstimator(\n",
" freq=dataset.metadata.freq,\n",
" prediction_length=dataset.metadata.prediction_length,\n",
" context_length=dataset.metadata.prediction_length*6,\n",
" \n",
" scaling=True,\n",
" num_feat_static_cat=len(dataset.metadata.feat_static_cat),\n",
" cardinality=[int(cat_feat_info.cardinality) for cat_feat_info in dataset.metadata.feat_static_cat],\n",
" embedding_dimension=[5],\n",
" \n",
" enc_dec_config=enc_dec_config,\n",
" \n",
" batch_size=256,\n",
" num_batches_per_epoch=100,\n",
" trainer_kwargs=dict(gpus=\"1\", max_epochs=50, logger=CSVLogger(\".\", \"lightning_logs/\")),\n",
" )"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "c77b420c",
"metadata": {
"scrolled": true
},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"/home/kashif/.env/pytorch/lib/python3.10/site-packages/pytorch_lightning/utilities/parsing.py:262: UserWarning: Attribute 'model' is an instance of `nn.Module` and is already saved during checkpointing. It is recommended to ignore them using `self.save_hyperparameters(ignore=['model'])`.\n",
" rank_zero_warn(\n",
"/home/kashif/.env/pytorch/lib/python3.10/site-packages/pytorch_lightning/trainer/connectors/accelerator_connector.py:446: LightningDeprecationWarning: Setting `Trainer(gpus='1')` is deprecated in v1.7 and will be removed in v2.0. Please use `Trainer(accelerator='gpu', devices='1')` instead.\n",
" rank_zero_deprecation(\n",
"GPU available: True (cuda), used: True\n",
"TPU available: False, using: 0 TPU cores\n",
"IPU available: False, using: 0 IPUs\n",
"HPU available: False, using: 0 HPUs\n",
"/home/kashif/.env/pytorch/lib/python3.10/site-packages/pytorch_lightning/utilities/parsing.py:101: UserWarning: attribute 'model' removed from hparams because it cannot be pickled\n",
" rank_zero_warn(f\"attribute '{k}' removed from hparams because it cannot be pickled\")\n",
"/home/kashif/.env/pytorch/lib/python3.10/site-packages/pytorch_lightning/trainer/configuration_validator.py:108: PossibleUserWarning: You defined a `validation_step` but have no `val_dataloader`. Skipping val loop.\n",
" rank_zero_warn(\n",
"LOCAL_RANK: 0 - CUDA_VISIBLE_DEVICES: [0]\n",
"\n",
" | Name | Type | Params\n",
"------------------------------------------\n",
"0 | model | TorchscaleModel | 165 K \n",
"------------------------------------------\n",
"165 K Trainable params\n",
"0 Non-trainable params\n",
"165 K Total params\n",
"0.661 Total estimated model params size (MB)\n"
]
},
{
"data": {
"application/vnd.jupyter.widget-view+json": {
"model_id": "093dab1c2d694f2b879622f8a21f139e",
"version_major": 2,
"version_minor": 0
},
"text/plain": [
"Training: 0it [00:00, ?it/s]"
]
},
"metadata": {},
"output_type": "display_data"
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"Epoch 0, global step 100: 'train_loss' reached 6.82088 (best 6.82088), saving model to './lightning_logs/version_10/checkpoints/epoch=0-step=100.ckpt' as top 1\n",
"Epoch 1, global step 200: 'train_loss' reached 5.84754 (best 5.84754), saving model to './lightning_logs/version_10/checkpoints/epoch=1-step=200.ckpt' as top 1\n",
"Epoch 2, global step 300: 'train_loss' reached 5.61930 (best 5.61930), saving model to './lightning_logs/version_10/checkpoints/epoch=2-step=300.ckpt' as top 1\n",
"Epoch 3, global step 400: 'train_loss' reached 5.47171 (best 5.47171), saving model to './lightning_logs/version_10/checkpoints/epoch=3-step=400.ckpt' as top 1\n",
"Epoch 4, global step 500: 'train_loss' reached 5.35223 (best 5.35223), saving model to './lightning_logs/version_10/checkpoints/epoch=4-step=500.ckpt' as top 1\n",
"Epoch 5, global step 600: 'train_loss' reached 5.30912 (best 5.30912), saving model to './lightning_logs/version_10/checkpoints/epoch=5-step=600.ckpt' as top 1\n",
"Epoch 6, global step 700: 'train_loss' reached 5.25711 (best 5.25711), saving model to './lightning_logs/version_10/checkpoints/epoch=6-step=700.ckpt' as top 1\n",
"Epoch 7, global step 800: 'train_loss' was not in top 1\n",
"Epoch 8, global step 900: 'train_loss' reached 5.22801 (best 5.22801), saving model to './lightning_logs/version_10/checkpoints/epoch=8-step=900.ckpt' as top 1\n",
"Epoch 9, global step 1000: 'train_loss' reached 5.16079 (best 5.16079), saving model to './lightning_logs/version_10/checkpoints/epoch=9-step=1000.ckpt' as top 1\n",
"Epoch 10, global step 1100: 'train_loss' was not in top 1\n"
]
}
],
"source": [
"predictor = estimator.train(\n",
" training_data=dataset.train,\n",
" shuffle_buffer_length=1024,\n",
" num_workers=8,\n",
" cache_data=True,\n",
" )"
]
},
{
"cell_type": "code",
"execution_count": 9,
"id": "f8a362b6",
"metadata": {},
"outputs": [],
"source": [
"forecast_it, ts_it = make_evaluation_predictions(\n",
" dataset=dataset.test,\n",
" predictor=predictor,\n",
")"
]
},
{
"cell_type": "code",
"execution_count": 10,
"id": "5fdc12da",
"metadata": {},
"outputs": [],
"source": [
"forecasts = list(forecast_it)"
]
},
{
"cell_type": "code",
"execution_count": 11,
"id": "4b7d3409",
"metadata": {},
"outputs": [],
"source": [
"tss = list(ts_it)"
]
},
{
"cell_type": "code",
"execution_count": 12,
"id": "9b154bde",
"metadata": {},
"outputs": [],
"source": [
"evaluator = Evaluator()"
]
},
{
"cell_type": "code",
"execution_count": 13,
"id": "0fdec8a7",
"metadata": {},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"\n",
"Running evaluation: 2247it [00:00, 3399.98it/s]\n",
"/home/kashif/.env/pytorch/lib/python3.10/site-packages/pandas/core/dtypes/astype.py:170: UserWarning: Warning: converting a masked element to nan.\n",
" return arr.astype(dtype, copy=True)\n"
]
}
],
"source": [
"agg_metrics, ts_metrics = evaluator(iter(tss), iter(forecasts))"
]
},
{
"cell_type": "code",
"execution_count": 14,
"id": "7f28f4d3",
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"{'MSE': 10018107.873432416,\n",
" 'abs_error': 26035099.499095917,\n",
" 'abs_target_sum': 128632956.0,\n",
" 'abs_target_mean': 2385.272140631954,\n",
" 'seasonal_error': 189.49338196116761,\n",
" 'MASE': 2.4701614804361784,\n",
" 'MAPE': 0.25477670689181103,\n",
" 'sMAPE': 0.24819572961811426,\n",
" 'MSIS': 20.209983280943216,\n",
" 'QuantileLoss[0.1]': 11107718.37581723,\n",
" 'Coverage[0.1]': 0.1886774959204866,\n",
" 'QuantileLoss[0.2]': 17534385.729109436,\n",
" 'Coverage[0.2]': 0.2506860999851654,\n",
" 'QuantileLoss[0.3]': 21799774.12036043,\n",
" 'Coverage[0.3]': 0.29637665034861294,\n",
" 'QuantileLoss[0.4]': 24542475.651297044,\n",
" 'Coverage[0.4]': 0.3366525738021065,\n",
" 'QuantileLoss[0.5]': 26035099.4930306,\n",
" 'Coverage[0.5]': 0.37475893784305,\n",
" 'QuantileLoss[0.6]': 26299950.11746931,\n",
" 'Coverage[0.6]': 0.4088785046728972,\n",
" 'QuantileLoss[0.7]': 24979193.32426689,\n",
" 'Coverage[0.7]': 0.4523809523809524,\n",
" 'QuantileLoss[0.8]': 21707479.345738567,\n",
" 'Coverage[0.8]': 0.5081590268506158,\n",
" 'QuantileLoss[0.9]': 15351255.969729993,\n",
" 'Coverage[0.9]': 0.6027110221035455,\n",
" 'RMSE': 3165.1394714028665,\n",
" 'NRMSE': 1.3269510918633773,\n",
" 'ND': 0.20239836126517932,\n",
" 'wQuantileLoss[0.1]': 0.08635204166354717,\n",
" 'wQuantileLoss[0.2]': 0.13631332338432334,\n",
" 'wQuantileLoss[0.3]': 0.16947269811913854,\n",
" 'wQuantileLoss[0.4]': 0.19079461760403799,\n",
" 'wQuantileLoss[0.5]': 0.2023983612180272,\n",
" 'wQuantileLoss[0.6]': 0.2044573252088626,\n",
" 'wQuantileLoss[0.7]': 0.1941896859174012,\n",
" 'wQuantileLoss[0.8]': 0.16875519323165183,\n",
" 'wQuantileLoss[0.9]': 0.11934154704281222,\n",
" 'mean_absolute_QuantileLoss': 21039703.56964661,\n",
" 'mean_wQuantileLoss': 0.1635638659322002,\n",
" 'MAE_Coverage': 0.15104954754487465,\n",
" 'OWA': nan}"
]
},
"execution_count": 14,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"agg_metrics"
]
},
{
"cell_type": "code",
"execution_count": 15,
"id": "cc3f804d",
"metadata": {},
"outputs": [
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAAB7kAAAXCCAYAAABuS1l6AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjYuMSwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/av/WaAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdeZhcdZk2/ruqurqq9y3pbJ2NBAIREgISCCCEgL4MiAKCv3GZSVB0GEdfGUEc3sFZUF98RVRmlJEZEVAHHIhsISZsSUggG9nXTjq971vt26nlnN8fp87p6u5au2s5VXV/rquvJLX1t6srtZznez+PTpIkCURERERERERERERERERERHlAn+sFEBERERERERERERERERERJYtFbiIiIiIiIiIiIiIiIiIiyhsschMRERERERERERERERERUd5gkZuIiIiIiIiIiIiIiIiIiPIGi9xERERERERERERERERERJQ3WOQmIiIiIiIiIiIiIiIiIqK8wSI3ERERERERERERERERERHlDRa5iYiIiIiIiIiIiIiIiIgob7DITUREREREREREREREREREeYNFbiIiIiIiIiIiIiIiIiIiyhsschMR5YDT6cS//Mu/4JJLLkFlZSVqampwxRVX4IknnoDf78/18oiIiIgoD3g8HmzZsgU//OEPceedd2LhwoXQ6XTQ6XT4l3/5l1wvj4iIiIjyxOjoKJ599ll8+ctfxvLly1FRUQGTyYSmpibcfvvtePXVV3O9RCKiSXSSJEm5XgQRUTHp7OzE2rVr0dHRAQAoLy9HKBSCIAgAgFWrVuG9995DXV1dDldJRERERFq3Y8cO3HDDDVHP++d//mcWuomIiIgoKUajEcFgUP232WyGwWCA2+1WT/uLv/gLbNy4EeXl5blYIhHRJExyExFlUTAYxG233YaOjg7MmTMH77zzDtxuNzweD/74xz+iqqoKhw8fxpe//OVcL5WIiIiI8kBdXR1uvPFGfPe738WLL76I2bNn53pJRERERJRngsEgVq9ejaeeegqtra3wer1wuVxob2/HV7/6VQDAli1b8Dd/8zc5XikR0RgmuYmIsuiZZ57BvffeCwDYvXs31qxZM+78F198EV/84hcBAO+++y5uvPHGrK+RiIiIiPJDKBSCwWAYd9qiRYvQ2dnJJDcRERERJW379u0xOwQBwH333Yenn34aANDV1YX58+dna2lERDExyU1ElEXPP/88AOCGG26YVOAGgL/8y7/E4sWLAQC/+93vsro2IiIiIsovEwvcRERERERTEa/ADUBNcwPAgQMHMr0cIqKksMhNRJQlHo8HH374IQB5hk00Op0ON998MwDg7bffztraiIiIiIiIiIiIiKIxm83q30OhUA5XQkQ0hkVuIqIsOX36NERRBABcfPHFMS+nnDcwMACLxZKVtRERERERERERERFFs2PHDvXvl1xySe4WQkQUgUVuIqIs6evrU/8+b968mJeLPC/yOkRERERERERERETZZLPZ8NhjjwEAPvGJT2DZsmU5XhERkYxFbiKiLHE6nerfy8vLY14u8rzI6xARERERERERERFliyiK+Ku/+iv09/fDbDbjl7/8Za6XRESkYpGbiIiIiIiIiIiIiIiIxvn2t7+NN998EwDwq1/9CitWrMjxioiIxrDITUSUJVVVVerfPR5PzMtFnhd5HSIiIiIiIiIiIqJsePDBB9Xk9s9//nN85StfyfGKiIjGY5GbiChL5s6dq/69t7c35uUiz4u8DhEREREREREREVGmPfTQQ3jiiScAAD/96U9x//3353ZBRERRsMhNRJQlF110EfR6+Wn3xIkTMS+nnDd79mzU19dnZW1ERERERERERERE3/3ud/H4448DAH7yk5/ggQceyPGKiIiiY5GbiChLysvLcc011wAAtm7dGvUykiThrbfeAgB86lOfytraiIiIiIiIiIiIqLg9+OCD+OlPfwpALnB/97vfzfGKiIhiY5GbiCiL1q9fDwDYvn079u3bN+n8l19+GW1tbQCAv/7rv87q2oiIiIiIiIiIiKg4Pfjgg+NalLPATURaxyI3EVEWrV+/HpdccgkkScLnPvc5vPfeewAAURTx8ssv42tf+xoA4C/+4i9w44035nKpRERERJQHrFYrRkZG1C9RFAEAHo9n3OkulyvHKyUiIiIirYqcwf2zn/2MLcqJKC/oJEmScr0IIqJi0tHRgRtuuAEdHR0A5DbmoijC5/MBAFatWoX33nsPdXV1OVwlEREREeWDRYsWobOzM+Hl1q9fj+eeey7zCyIiIiKivNLV1YWFCxcCAPR6PWbOnBn38g8++CAefPDBbCyNiCiuklwvgIio2CxatAjHjh3DT3/6U7zyyitob2+H0WjExz72MXzhC1/At771LZSWluZ6mURERERERERERFTglE5Ayt8HBwfjXp4dgohIK5jkJiIiIiIiIiIiIiIiIiKivMGZ3ERERERERERERERERERElDdY5CYiIiIiIiIiIiIiIiIiorzBIjcREREREREREREREREREeUNFrmJiIiIiIiIiIiIiIiIiChvsMhNRERERERERERERERERER5oyTXC9ACURTR19eHqqoq6HS6XC+HiIiIKCskSYLT6cTcuXOh13PvYzrwfSUREREVI76vTD++ryQiIqJilMr7Sha5AfT19WH+/Pm5XgYRERFRTnR3d6OpqSnXyygIfF9JRERExYzvK9OH7yuJiIiomCXzvpJFbgBVVVUA5Dusuro6x6shIiIiyg6Hw4H58+er74Vo+vi+koiIiIoR31emH99XEhERUTFK5X0li9yA2vKnurqabxqJiIio6LD9YfrwfSUREREVM76vTB++ryQiIqJilsz7Sg7JISIiIiIiIiIiIiIiIiKivMEiNxERERERERERERERERER5Q0WuYmIiIiIiIiIiIiIiIiIKG+wyE1ERERERERERERERERERHmDRW4iIiIiIiIiIiIiIiIiIsobLHITEREREREREREREREREVHeYJGbiIiIiIiIiIiIKMLBgwfx4x//GHfeeSeampqg0+mg0+lSuo2bbrpJvV5PT0+GVkpERERUnEpyvQAiIiIiIiIiIiIiLfnBD36A119/fcrXf+655/Dee+9Bp9NBkqQ0royIiIiIABa5iYiIiIiIiIiIiMZZs2YNVqxYgSuuuAJXXHEFFi1aBEEQkrru8PAwHnjgAXzqU5/CmTNn0NnZmeHVEhERERUfFrmJiIiIiIiIiIiIInzve9+b8nXvv/9+eDwePPXUU7jxxhvTuCoiIiIiUnAmNxEREREREREREVEabN26FS+88AL+8R//EUuWLMn1coiIiIgKFovcRERERERERERERNPkdrvxt3/7t7jwwgvx0EMP5Xo5RERERAWNRW4iIiIiIiIiymtbjvfjuy8fhRAM5XopRFTE/umf/gkdHR349a9/jdLS0pSuKwgCHA7HuC/KvWM9NvzvFw+jx+rJ9VKIiIhoAha5iYiIiJIQFIO5XgIRERHF8O/bzuHlgz3Y327J9VKIqEgdOnQITz75JNavX4/rr78+5es/9thjqKmpUb/mz5+fgVVSqn63pxNvHO3DG0f7cr0UIiIimoBFbiIiIqIkWLw8aE5ERKRVvoCc4LZ5AjleCREVo1AohHvvvRe1tbX46U9/OqXbePjhh2G329Wv7u7uNK+SpkJ5XXH6uOmZiIhIa0pyvQAiIiKifGDxWtBY0ZjrZRAREVEUQlAEADh8LHITUfb94he/wOHDh/HMM89gxowZU7oNk8kEk8mU5pXRdCmvK26BRW4iIiKtYZGbiIiIKAFREmH32XO9DCIiIorBH5KL3EzaEVEubNq0CTqdDs8//zx+97vfjTtvYGAAAHD33XfDZDLhH/7hH3DzzTfnYpk0BcrrilsI5XglRERENBGL3EREREQJCEEB7oA718sgIiKiGPxKktvLJDcR5YYkSdi5c2fM8/fu3QsA2LBhQ5ZWROmgvK54/NxERUREpDUschMREREl4Av64Al4cr0MIiIiikEpcjPJTUS5sGPHjpjnLVq0CJ2dneju7kZTU1P2FkVp4Qy3K3exXTkREZHm6HO9ACIiIiKtE0IC3H4muYmIiLRKaVfOmdxERJQuoijBGS5ue/xsV05ERKQ1THITERERJcAkNxERkXaFRAkhUQLAduVElD6bN2/GD37wA/Xffr8fAHDVVVepp33/+9/HrbfemvW1UXa4/UFIUvjvTHITERFpDovcRERERAkIQQFCSEBIDMGgN+R6OURERBRBaVUOsF05EaXP8PAw9u3bN+n0yNOGh4ezuSTKMkfEa4qbM7mJiIg0h0VuIiIiogSEkAAA8AQ8qDJV5Xg1REREFCmyyM125USULhs2bMCGDRumfTsdHR3Tvg3KDWfEa4pHYLtyIiIireFMbiIiojzjD/lzvYSi4wv6AADuAOdyExERaY0QGis8MMlNRETp4vCOvaa42K6ciIhIc1jkJiIiyjNt1rZcL6HoqEVuP4vcREREWhMISerfOZObiIjSJTLJLQRFBENinEsTERFRtrHITURElGfOjp7N9RKKjhAca1dORERE2hLZrtztD7EIQUREaTFxBIbbz5blREREWsIiNxERUZ4Zdg/D6rXmehlFJXImNxEREWlLZJEbYEtZIiJKj4kjMDx+vr4QERFpCYvcREREeSQoBhEQA+hx9OR6KUWFM7mJiIi0a2KRm3O5iYgoHSaOwHALTHITERFpCYvcREREeURJErPInV1sV05ERKRd/tD4ooOdc7mJiCgNJm6acrNTCBERkaawyE1ERJRHlCJrv6sfkiTleDXFQ01y+5nkJiIi0hqBSW4iIsqAyTO5+fpCRESkJSxyExER5RGlyOoP+THkHsrxaoqDKIkIiPLBDSa5iYiItGdiu/KJRQkiIqKpcExKcrNdORERkZawyE1ERJRHIousbFmLine truncated
"text/plain": [
"<Figure size 2000x1500 with 9 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"plt.figure(figsize=(20, 15))\n",
"date_formater = mdates.DateFormatter('%b, %d')\n",
"plt.rcParams.update({'font.size': 15})\n",
"\n",
"for idx, (forecast, ts) in islice(enumerate(zip(forecasts, tss)), 9):\n",
" ax = plt.subplot(3, 3, idx+1)\n",
"\n",
" plt.plot(ts[-4 * dataset.metadata.prediction_length:].to_timestamp(), label=\"target\", )\n",
" forecast.plot( color='g')\n",
" plt.xticks(rotation=60)\n",
" ax.xaxis.set_major_formatter(date_formater)\n",
" ax.set_title(forecast.item_id)\n",
"\n",
"plt.gcf().tight_layout()\n",
"plt.legend()\n",
"plt.show()"
]
},
{
"cell_type": "code",
"execution_count": 18,
"id": "6f03bfd2",
"metadata": {},
"outputs": [],
"source": [
"metrics = pd.read_csv(\"lightning_logs/version_86/metrics.csv\")"
]
},
{
"cell_type": "code",
"execution_count": 19,
"id": "8e76b769",
"metadata": {},
"outputs": [
{
"data": {
"text/html": [
"<div>\n",
"<style scoped>\n",
" .dataframe tbody tr th:only-of-type {\n",
" vertical-align: middle;\n",
" }\n",
"\n",
" .dataframe tbody tr th {\n",
" vertical-align: top;\n",
" }\n",
"\n",
" .dataframe thead th {\n",
" text-align: right;\n",
" }\n",
"</style>\n",
"<table border=\"1\" class=\"dataframe\">\n",
" <thead>\n",
" <tr style=\"text-align: right;\">\n",
" <th></th>\n",
" <th>train_perplexity</th>\n",
" <th>epoch</th>\n",
" <th>step</th>\n",
" <th>val_loss</th>\n",
" <th>train_loss</th>\n",
" </tr>\n",
" </thead>\n",
" <tbody>\n",
" <tr>\n",
" <th>0</th>\n",
" <td>2.042362</td>\n",
" <td>0</td>\n",
" <td>49</td>\n",
" <td>NaN</td>\n",
" <td>NaN</td>\n",
" </tr>\n",
" <tr>\n",
" <th>1</th>\n",
" <td>2.050069</td>\n",
" <td>0</td>\n",
" <td>99</td>\n",
" <td>NaN</td>\n",
" <td>NaN</td>\n",
" </tr>\n",
" <tr>\n",
" <th>2</th>\n",
" <td>2.743227</td>\n",
" <td>0</td>\n",
" <td>149</td>\n",
" <td>NaN</td>\n",
" <td>NaN</td>\n",
" </tr>\n",
" <tr>\n",
" <th>3</th>\n",
" <td>2.440984</td>\n",
" <td>0</td>\n",
" <td>199</td>\n",
" <td>NaN</td>\n",
" <td>NaN</td>\n",
" </tr>\n",
" <tr>\n",
" <th>4</th>\n",
" <td>NaN</td>\n",
" <td>0</td>\n",
" <td>199</td>\n",
" <td>4.355846</td>\n",
" <td>NaN</td>\n",
" </tr>\n",
" <tr>\n",
" <th>...</th>\n",
" <td>...</td>\n",
" <td>...</td>\n",
" <td>...</td>\n",
" <td>...</td>\n",
" <td>...</td>\n",
" </tr>\n",
" <tr>\n",
" <th>295</th>\n",
" <td>80.659866</td>\n",
" <td>49</td>\n",
" <td>9899</td>\n",
" <td>NaN</td>\n",
" <td>NaN</td>\n",
" </tr>\n",
" <tr>\n",
" <th>296</th>\n",
" <td>82.568138</td>\n",
" <td>49</td>\n",
" <td>9949</td>\n",
" <td>NaN</td>\n",
" <td>NaN</td>\n",
" </tr>\n",
" <tr>\n",
" <th>297</th>\n",
" <td>81.211136</td>\n",
" <td>49</td>\n",
" <td>9999</td>\n",
" <td>NaN</td>\n",
" <td>NaN</td>\n",
" </tr>\n",
" <tr>\n",
" <th>298</th>\n",
" <td>NaN</td>\n",
" <td>49</td>\n",
" <td>9999</td>\n",
" <td>1.084462</td>\n",
" <td>NaN</td>\n",
" </tr>\n",
" <tr>\n",
" <th>299</th>\n",
" <td>NaN</td>\n",
" <td>49</td>\n",
" <td>9999</td>\n",
" <td>NaN</td>\n",
" <td>1.707654</td>\n",
" </tr>\n",
" </tbody>\n",
"</table>\n",
"<p>300 rows × 5 columns</p>\n",
"</div>"
],
"text/plain": [
" train_perplexity epoch step val_loss train_loss\n",
"0 2.042362 0 49 NaN NaN\n",
"1 2.050069 0 99 NaN NaN\n",
"2 2.743227 0 149 NaN NaN\n",
"3 2.440984 0 199 NaN NaN\n",
"4 NaN 0 199 4.355846 NaN\n",
".. ... ... ... ... ...\n",
"295 80.659866 49 9899 NaN NaN\n",
"296 82.568138 49 9949 NaN NaN\n",
"297 81.211136 49 9999 NaN NaN\n",
"298 NaN 49 9999 1.084462 NaN\n",
"299 NaN 49 9999 NaN 1.707654\n",
"\n",
"[300 rows x 5 columns]"
]
},
"execution_count": 19,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"metrics"
]
},
{
"cell_type": "code",
"execution_count": 20,
"id": "ad490889",
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"Text(0, 0.5, 'perplexity')"
]
},
"execution_count": 20,
"metadata": {},
"output_type": "execute_result"
},
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYwAAAEQCAYAAACjnUNyAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjUuMSwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/YYfK9AAAACXBIWXMAAAsTAAALEwEAmpwYAABKoUlEQVR4nO2dd3ib5dW47yPJe8WOnb0XIWQBIewZ9mjpoIP2a4G20E1p+5VCy9fyle71a6FfKV1QKKVAgZYNYUNYCSRkbzvO8N6WbVnS8/vjfaXI8pJsLUvnvi5dtt6l82g85z3PWWKMQVEURVGGw5FsARRFUZSxgSoMRVEUJSJUYSiKoigRoQpDURRFiQhVGIqiKEpEuJItQKwpLy83s2bNSrYYiqIoY4p169Y1GGMqhjom7RTGrFmzWLt2bbLFUBRFGVOISNVwx+iSlKIoihIRqjAURVGUiFCFoSiKokSEKgxFURQlIlRhKIqiKBGhCkNRFEWJCFUYiqIoSkSowlAUJe3p9fm5540qujy+ZIsypkm7xD1FUTIDv9/Q1t3LuPzsAfcbY3h84yEeefcgR04u4tbnd9Hd6+Ozp85JsKTpgyoMRVHGJL9/aTc/f3o7k4pzOXrGOE6aO57l00u5/eXdnH/UJO5cU8m6qmYAVm+tBeD+tdV85pTZiEifazV3eijIcZHt0kWXoVCFoSjKmOJQaxdv7W3ir69VctSUYuZWFLK+uoUnN9WQ5RR6fYbH3ztEeWEOP/3QEtweHzc/uoUT5pTxxp4m1lU1s2JWWfB6HT1eTv3ZCwCcvqCCVUdO4OKlU1R5DIAqDEVRxhS/fnYH96/db/3/0WWcOt+ql3frczv5y2t7+X8fO5q6tm4uWDKZwhwXxhhOmDOe6WX5nPHzF7nl8a089IWTcDgsK2NDdQsdPV5OW1DBW5VNPL7xENtq2rnxwiOTNsZURVWooihjih21HQB87tTZnDKvPLj9K6vms+6753D6ggouWzGdwhzrflhEOHJyMYU5Lm68cCHrq1t4YF118Lx1Vc2IwK0fP5o3b1jFmUdU8Ph7hzDG9HndP768h4fe2Z+AEaYuqjCUIL0+P70+f7LFUJRB8fkN22vaufLkWXznokX9fBEBq2EwPnD0VI6bVcpPn9pOi9sDwDv7mlkwoYiSvCwcDuH8xZM40NLFlkNtwfMOtHTxk6e28aMntuGN8Deyq66D57fVRjnC1EYVhhLkh49v5bgfrualHfXJFkVRBqSysZOuXh+LJheP6HwR4eb3LabF7eEXz2zH7ze8U9XMMTPHBY85a+FEROC6f65nza4GAO5aU4nPb2jo6OHlnUP/Powx3L+2motvfYXP3LWW9u7efsf0eH2s2d0wojEkE1UYSpA39jTS4u7lir++xW9W78TvN8OfpChxxO83/PPtfXzn4Y384619PL25BoBFU0amMALnfurEWfz9zX386dU9tHV7OXHu4aWtiqIcfvSBJdS19/CX1yrp7PHyj7f2cd5RExlfkM2/3jnA5oOt3L+2mnVVTTR3eoLnuj1evvHABr714HuUF+ZgDGyvae/z+sYY/uvPb3H5H99k4/7WEY8jGajTWwHA4/Wzu76DK06aRVtXL79evYNd9R3c+vGjky2aksH8evUObn1+F3lZTv7+5j4AspzCvAmFo7rudecs4LH3DvKjJ7ZRnOvi3EUT++z/+MoZvL23iZd3NvDA2mrau71cc/pcxuVV88TGQ2w52Mbehs7g8ecumsgdn1rBz57azsPvHuDaVfO5bMU0TvnpC2w51NYnKuv+tdW8tbcJgNf3NLBkWklEMr+wrY7nt9XxkRXTIz4n1qiFoQCwp6GDXp/h6Bnj+OVHlnHlybN4dMPB4DqvEjtqWrv517r91Lf38KdX9tDa1X/JQrF4Y08jy6ePY/PN5/HiN8/gNx9bzu2fPJYcl3NU1y3Jy+L68xcCcOnRU8nN6n+9ZdPH0dDRw20v7OKYGeM4ZkYpZxxRQXuPl70NnXzjnAX85YoVnLVwAi9ur8fj9fPC9jrOOmIC152zgKnj8hiXn8XWEF+I2+Pll8/s4JgZ45hdXhBUHAH8fsPdr1fyrQc3UNvWHdze3Onhqrve5u43qvjqfe/S3ZucjHVVGAoA2w5ZZvORk4sREc6x77huf2kPNzy0EY9XneGx4pbHt/CNBzZw7X3vcsvjWzn7Vy/x2HsH2dvQyQ8e20JNazcPrttPXXv38BdLc/Y3dzGnogCHQ5hVXsD7l09l1ZEThz8xAj50zDRuuXQxXz5r3oD7l00fB0BDh4fPnGJlh588vxynQ3AIfPz4GZy1cCIfOHoqHp+fF7fXUdXo5iQ7cktEWDS5mC0HDyuMP7+yl7r2Hm688EiOn13GW3ub8IUs/T6x6RA3/Xsz96/dzzNbDjvMt9a0YQx8/vS57G3o5P9e2BWT9yBadElKAWBbTTtZTmF2eQEAy6aNwyFw+0u7ASjIdvLdixclU8S0YH+zmyc3Wevwa3Y3smJmKT1eP1++912yXQ48Xj/3vrmPrl4f2U4HHzxmKp89dTbzJhQlWfL+VDV2cs3d6/jjp1YwvSw/5tfv8fqoaetmemnsrw1WRNUnT5g56P4jJxeR5RQmFOVy3lGWkirOzeI0W2mUF+YAsGSqtTx0x8t7ADh53viQaxRzzxtV9Pr8tHb1cvtLuznvqImsmFXGviY3971dzc66dhZOsnwy9765j6nj8mh2e9hd1xG8TsAPctUpszjY0sUfXt7DR46bzjT7valr6yY/xxUMJY4XamEoAGw80MKCiUVkOa2vREGOK/glnlaax59e3cu6qr7m83v7Wzj1Z8/rslWEbDnYxlV3vo1DDk8yV50ym0e+dDLfu2QRS6aW8KUz5+L1+7nhgoV85LhpPPzuAc799cu8t78lucIPQCDBbbiooZFysKUbY4iLMoqEHJeTb5x7BP/7/qNwOQ9PlXd8agW//+SxweczyvLJz3aytqqZiqIcjph4WLkHbgje29/Cb5/bSbfXz7fspbDF9ncgsGS15WAba3Y3cvnxM5g/oZBdYQqjrCCbisIcvn3BQkTgx09uAywn+odvf50bHtoYvzfDJuEKQ0Q+JiLviEiHiBwQkb+JyJSwY0REbhSRahHpEpGXRWR5omXNFDxeP+uqmlk5u6zP9pWzy8hxOfjnNScypSSXGx/a1CdP4739rVQ3dfX5YsebZzbX8KdX9iTs9UZKTWt30Cnq8xv+8NJu3v+7V2lx9/LnTx/H9ecv5IQ5ZZy1cAJOh3DlybP51xdO4r/PW8jG75/HNafP5ZZLl/D8N8/Ab+D13Y0JH8OBli7uWlNJe3cvT27sn8j2sh1+velAfCJ9qpvcAEwvzYvL9SPh86fP7bcEluV0BG+swLJUnHY+yHcuPLJPbsiJc8cjAve8sY9739zHx46bztwKy2E/u7yAbKeDbTXteLx+/vvBDZQXZnP5yhnMreirMLbVtHPExCJEhCnj8vj86XN5/L1DvLW3if3NXexrcrN6Sy1ujzeeb0diFYaIvA/4B7AGeD9wPXAa8LiIhMrybeAm4KfAJUAHsFpEJiVS3kzhvf0tdPf6OX72+D7brztnAf/58ilMHZfHze9fzPbadv70yt7g/oBlUdfekzBZ71+7n9uStH4bKf9ef4Czf/USn/7LWwD8/Ont/PjJbaxaOJGnv3Yapy2o4JT55dx39YkDOltDt00dl8fE4px+oZmJ4Lerd/K9/2zmF09v5wt/f4c1IUqrs8cbLOy36UDbYJcYFdXNtsJIkoURDb/+6HK+fs4C3r+8z70v4/KzOWJiEQ+/e4DcLCfXnj0/uC/L6WDuhEK217Rz2wu72HywjR9+YAmlBdnMnVBITVs37d297Gt0s6O2nSMmHbZcrjltLlNKcvnfxzYHHeddvT5e2h7fHKpEWxiXA+8YY75sjHnOGHMP8FVgOXAEgIjkYimMHxtjbjPGrAYuAwzw5QTLmxG8sceaCI4PszBK8rKCX9JzFk3kvKMm8pvndgTv/JrdVnRPXVvinLOtXR5a3L0puQzm9ni5/sH3uPa+9Rhj2Nfk5kBLF397vZJLlk3h9588htKCgUtxD8URk4rZXptYhdHd6+OJTYcA+PeGgwD8+dXDNwuv7mqg12dYOq2E7fYdcqypbuoiyylMLM6N+bVjzdmLJvLVVfP7ZZ4DHD2jFIAffXAJE4r6jmXhpCLWVjbzuxd28cGjp3LeUdY9cSBs+JN/fovTf/ECPV5/n9DfvGwn11+wkE0H2vjGAxsoynExviCbJ2z/WLxItMLIAsLt1xb7b+CdPgkoBu4PHGCM6QQeBS6Is3wZR1t3L/9cW81RU4qHncy+/76jcIpw0783YYwJJiwl0sJosZVUZaM77q/1/LZa/r3+AK/vbmRPfcegiYyt7l5au3q5/l8buX9dNV86cy4//fBSAG55bAtuj48vnTl3wMkkEhZOKmJnXUfEJSliwYvb62jvtpY3Wty9OASe31bHLY9tocvj48mNhxiXn8VVJ8/G4/OzIw4KbX+zm6nj8nAOU+4j1fnWeUfw1yuO433LpvTbt3BSER09XioKc/jeJUcFty+aXIzTIRxs6eKrZ81nzbfPCkZfBXjfsimcfeQEAOZOKOQHly7ms6fMjutYEh0l9RfgERH5FPAIMAm4BXjeGLPFPmYh4AN2hp27FfhoguTMGH78xDYOtnTz66uXD3vs5JI8rjtnAbc8vpV1Vc00J2FJqsXOWahs6GS5HfYYD9bsauCqO9f22faFM+YGY/dDueLOtzAGdtS2c/nKGfz3eQvZaU+gT22uYdn0ccEAgpGwYGIRHq+fykb3qBPWIuXhdw9QXphDQY6TqkY3nzllNl29Pv706l6e2VJLY0cPFy2dzCw7qs4KAY5tMll9ew8TxoB1MRylBdmcuXDCgPtWzCrD5RB++uGllORnBbdPL8vnxW+ewcTi3EHLrIsIt3/yWO5cU8nLine truncated
"text/plain": [
"<Figure size 432x288 with 1 Axes>"
]
},
"metadata": {
"needs_background": "light"
},
"output_type": "display_data"
}
],
"source": [
"ax = metrics.train_perplexity.dropna().plot(kind=\"line\")\n",
"ax.set_xlabel(\"training steps\")\n",
"ax.set_ylabel(\"perplexity\")"
]
},
{
"cell_type": "code",
"execution_count": 21,
"id": "f1185a0f",
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"Text(0, 0.5, 'val neg. log likelihood')"
]
},
"execution_count": 21,
"metadata": {},
"output_type": "execute_result"
},
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAYEAAAEQCAYAAABWY8jCAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjUuMSwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/YYfK9AAAACXBIWXMAAAsTAAALEwEAmpwYAABPDUlEQVR4nO2dd3gc9bWw37Mqq96bbcmSbYrBBlcIJISa0FIgCemNmwRSSM9NclMuIZDce0O+m05C6oUUQhJIIxA6hF6NbWxjMO5N1ZZWbaWV9vf9MTPr1WrLrLRN0nmfZx/tzszOnNnVzpnTxRiDoiiKMjfxZFsARVEUJXuoElAURZnDqBJQFEWZw6gSUBRFmcOoElAURZnD5GdbgGSoq6szbW1t2RZDURRlRvHss892G2Pqo62bUUqgra2NZ555JttiKIqizChEZHesdeoOUhRFmcOoElAURZnDqBJQFEWZw6gSUBRFmcOoElAURZnDqBJQFEWZw6gSUBRFmcPMCSXwl+f28bsnY6bJKoqizFnmhBL4x4aD/O6JPdkWQ1EUJeeYE0qgrCifwdGxbIuhKIqSc8wJJVDqzWfAr0pAURQlkjmhBMq9+QyMqBJQFEWJZE4ogVJvPiNjQQLjwWyLoiiKklPMCSVQ5rWapQ6qNaAoijKBOaUE+jUuoCiKMoE5oQRKHUtAM4QURVEmMCeUQFmRuoMURVGiMTeUgDcPUHeQoihKJHNECRQAMDgynmVJFEVRcos5oQRKbUtgYCSQZUkURVFyizmhBMptS2BALQFFUZQJzAklELIENCagKIoygTmhBPLzPBQVeDRFVFEUJYI5oQTAKhjT7CBFUZSJZFUJiMgCERkQESMiZek8Vpk3X+sEFEVRIsi2JfBtYCATByrVTqKKoiiTyJoSEJHTgfOB/5eJ46kSUBRFmUxWlICI5AE/BK4GujNxzHIdLKMoijKJ/FgrRGQnYNzuyBizOInjfgTwAtcB707ifVOm1KsjJhVFUSKJqQSAW5moBN4BlAD3AJ1AA/BaYBC42e0BRaQWuAZ4jzEmICKJtr8cuBxg4cKFbg8zibIiDQwriqJEElMJGGP+3XkuIl8GtgOvM8YMhi0vA/4B+JI45jeBJ4wxd7jZ2BjzM+BnAGvXrnVtmUSiKaKKoiiTcRsTuAL4drgCADDGDGAFdq9wsxMRWQZ8ALhaRKpEpArLugCoFJFil/IkTZmOmFQURZlEPHdQOBVAY4x1TYDbHP+jgQLg8Sjr9gG/BD7kcl9JURo2YrKqpDAdh1AURZlxuFUCtwHfFhEf8HdjzKiIFAIXAd+y17vhEeCsiGXnA18ELgR2uNxP0pTbSmBAlYCiKEoIt0rgo8ANwB8BIyL9QDkgwN/t9QkxxnQDD4YvE5E2++nDtnspLZSGKQFFURTFwpUSMMb0AW+yffonYbmG2oGnjTFb0ihfytARk4qiKJNxawkAYIzZDGxOpQDGmBuwrIy0oiMmFUVRJuNaCdiZPB8GTgNqgEPAw8DPjDG96RAuleiISUVRlMm4ShEVkSXA81htHkqBPfbfq4GN9vqcRkdMKoqiTMatJfBdoBc4xRiz31koIguAO4DvYGUK5SxlocCwWgKKoigObovFzgSuDFcAAPbrq5mc9plzhLKDNCagKIoSwq0SMEBenH1MuZ1DpijI8+DN1xGTiqIo4bhVAg8A14hIa/hC+/XVwH2pFiwdlBfpTAFFUZRw3MYEPg3cD2wTkXVAB1YX0TXAXuCzaZEuxZTqTAFFUZQJuLIEjDG7gKXAJ7HqBAqALcDHgePs9TmPzhlWFEWZiOs6AWPMKHC9/ZiRlHrz6VcloCiKEiKpimEReQURxWLGmKfSIVg6KPfm0+7zZ1sMRVGUnMGVEhCRUuBPWB0/x4AeoBbIE5E7gbcaY4bSJmWK0GHziqIoE3GbHXQtcCrwdqDIGDMPKMIaOXkqVjvpnEdHTCqKokzErRJ4C/BFY8yfjDFBAGNM0BjzJ+A/gLemS8BUoiMmFUVRJuJWCVRipYJGYy/W5LGcp7RQR0wqiqKE41YJbAA+KiISvtB+/VF7fc6jMwUURVEm4jY76MvAP4GtIvIXjhSLvQloAy5Ii3QppizUSVRHTCqKooD7yWL3i8hq4D+x/P/zgIPAk8CbZ8x0MZ0poCiKMoFkisU2Y2UDzVh0poCiKMpE3MYEZgXlRTpTQFEUJZxkxkteArwZaMaqEZiAMebkFMqVFnSmgKIoykTcVgxfBVyJlQW0BRhNo0xpw5kuptlBiqIoFm4tgQ8C/2OM+XI6hUk3jhLQJnKKoigWbmMC5cyQwTHxKFVLQFEUZQJulcDNWM3jZjTOiMlcbiLX1T/CRdc9yt5DOd+PT1GUWUBMd5CIXBj28l7gWhGpA+4BeiO3N8bckXLp0kCuj5h84aCPDXt7eXrXIVpqSrItjqIos5x4MYF/YA2QD28V0Qa8P8q28QbR5xS5PmKyb9iqYdh3eDjLkiiKMheIpwQWZUyKDFJamNvtpHttJaDuIEVRMkFMJWCM2Z1JQTJFWVFuj5j0qSWgKEoGiRcTKHGmhYlIQuf0TJgsBlaaaEcOj5jsHbJKMPb1zoiPU1GUGU48d1C/iJxqzxAewPL7x2NGxATKvPnsyGFLwIkJHOj1MzYeJD9vTnX2UBQlw8RTAh8Atoc9T6QEZgS5PmfYUQLjQUO7z09ztWYIKYqSPuLFBG4Me35DRqTJALmeIto7FMAjEDRWXECVgKIo6WTO+RpKC/PxB4KM5eiIyb7hAEc1lAGaIaQoSvqJFxh+miRcQDOhiyiEj5gcp7Ik93Rg33CAkxfVsK1zQDOEFEVJO/FiApuZJXGAcJwRk/0jASpLCrIszWT6hgPUl3lpqihi72G1BBRFSS/xYgKXZlCOjJHLIyZHx4IMjY5TWVxAc3WxWgKKoqSdpPwhYtEiIq8UkdJ0CZVOcnnEpJMZVFlSQEt1CftVCSiKkmZcKwER+RiwH9gNPAwcay//s4h8Oi3SpYFcHjEZUgK2JXCwb5hAjgawFUWZHbhSAiLyeeA7wM+Bs5nYVO5B4O0plyxN5PKIyQlKoKaEoIGDvblb3awoyszH7WSxK4ArjTHXikhkZfCLwDGpFSt9lBbm7mCZvmGrZURlcQGF+ZZ+3nd4iIW1WiugKEp6cKsEmoBnY6wLEmXwfK7iuINysYmcYwlUlRRS57GMLc0QUhQlnbiNCbwMnBFj3elYw+dnBLk8YrJ36Ig7aF5lEXke0QwhRVHSiltL4HvAj0VkFLjFXtYgIh8EPgtclgbZ0oIzYjIXlYBjCVQU5ZOf57FqBbRqWFGUNOJKCRhjfiEi1cCVwNftxXcAQ8BVxpib0iRfWijz5uZMgb7hAOXe/FDnUK0VUBQl3bi1BDDGfFtErgdOBeqAQ8Djxpg+Eck3xuTeVTUGZUW5OV2sbyhARfGRKuaWmhIe2dadRYkURZntuE0R/QaAMabfGHO3MeYmY8ydtgIoBv6eVilTTGlhbs4Z7hsOUBXWyqK5upiOfj8jY7lX06AoyuzAbWD4kyLy5ciFIlIG3Akc72YnInKJiDwmIj0i4heRF0XkqyJSmITM06YsR9tJ9w4HqAy3BKpLMMYaMDOX+ctz+9i4rzfbYijKrMStErgI+IqIfMZZYMcIHgDmY2UIuaEWuB/4EHAB8CvgK1iFaBmjLEcHy0SzBMCqFZjLfP22Lfz0oR3ZFkNRZiVuA8MPiMibgb+KyDDwV+Aee/WrjTHtLvfz04hFD4hIBXCFiHzCGJORrqW5OmKyL8ISaK6xisT2Hpq7weFg0OAbDrC7ZzDboijKrMR17yBjzF1Y7SG+h1U4Ngyc7lYBxKEHyKg7KBdHTBpjJgWGmyqKyPfInLYEBkfHCBrY3T1Ehu4RFGVOEW+ozIVRFo8BNwFvwHLhnCpiVbYaY+5we1C79YQXWA18EvhJpqwAsGYK5JoS8AeCjI4HqSo+og/zPML8qmL2zuE0Uad2on9kjJ7BUerKvFmWSFFmF/HcQf/AGiojMdaH1wYYILKnUDwGsZQAwK+Bz8faUEQuBy4HWLhwYRKHiE2ZtyA0YtLJyc824c3jwrFqBeauJeAbPqKsd/cMqhJQlBQTTwksSuNxXwmUACdjFaD9CPhYtA2NMT8Dfgawdu3alFgLzkyBXBox2RvWPC6cluoS7n+xMxsi5QQ+/5G5D7u6h1jTWpNFaRRl9hFvstjudB3UGLPOfvqIiHQDN4rI/xpjtqfrmOGEZgqMjuXMiMm+Iad53GRLoKt/BH9gnKKCZIyt1GGM4eu3beG1xzfyqqPqMnpsx0IC2KXBYUVJOTFvg0WkJPx5osc0ZHAUQjotjwnk4kyB3ljuoBonTTR7cYFndx/mhsd2ceu6fRk/ts/+XArzPOzqmbtuMUVJF/F8If0icrL9fADoT/CYKq+y/+6cxj6SosxRAjkUHI4VE2iptvRrNuMCv33CMgq3d2X+Ttz5XI6bV86ubrUEFCXVxIsJfADYHvZ82v54EbkTuBfYDIxjKYDPAX/IlCsIclMJ+MLmC4fTbCuBbGUI9QyMcMfz7XgEdnQNYIzByQjLBD7/GCKwfEElf99wIOPHV5TZTryYwI1hz29I0fGeBi4F2rDSTXcAXwKuT9H+XVFWlHszBXqHAngEygonfiUN5V4K8zxZswT+9Ow+RseDvPPkhfz+qT10D4xSX565DB3fcIAybz6L68vo949xeChATWlGy0oUZVaT0dQYY8x/GmOWG2PKjDFVxpjVxpgfGmMCid+dOpwRk7kUE3CqhT2eiXe5Ho+woLqYfVmoGg4GDTc9uYeTF9Vw/vImALZ3DWRUBt9wgIqiAtrsEZs71SWkKCklXrHY0yThAjLGnJx4q9wglB2US5ZARMuIcLJVK/Dwy93sOTTEv593LIvrSgHY0TXIKYtrMyaDz299Lm328Xf3DLKmtTpjx1eU2U68mMBmUhAHyEVKczAm0DcLine truncated
"text/plain": [
"<Figure size 432x288 with 1 Axes>"
]
},
"metadata": {
"needs_background": "light"
},
"output_type": "display_data"
}
],
"source": [
"ax = metrics.val_loss.dropna().plot()\n",
"ax.set_xlabel(\"training steps\")\n",
"ax.set_ylabel(\"val neg. log likelihood\")"
]
},
{
"cell_type": "code",
"execution_count": 22,
"id": "d887cb3b",
"metadata": {},
"outputs": [],
"source": [
"X = predictor.prediction_net.vq_vae.embed.cpu()"
]
},
{
"cell_type": "code",
"execution_count": 23,
"id": "ae16d4bd",
"metadata": {},
"outputs": [],
"source": [
"X_embedded = TSNE(n_components=2, learning_rate='auto', init='random').fit_transform(X)"
]
},
{
"cell_type": "code",
"execution_count": 24,
"id": "8feaef88",
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"<matplotlib.collections.PathCollection at 0x7f4eec421b20>"
]
},
"execution_count": 24,
"metadata": {},
"output_type": "execute_result"
},
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAZIAAAECCAYAAADU5FG5AAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjUuMSwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/YYfK9AAAACXBIWXMAAAsTAAALEwEAmpwYAAAj/ElEQVR4nO3de7RcZZnn8e/PQGugJYkaR4nGoLYy9nhBz6hMVDQ6Ktqj8YLo0lHHS1rXrKHbCxoVR9GeMQ4NeJtZdMRLa2sztkB6kI6xJUILLWAyjJfGYIsRNGgLmohA5PrMH1WHVCpVp3bt67t3/T5rnXWSXbWr3jrJ2c9+n+d591ZEYGZmltc9mh6AmZm1mwOJmZkV4kBiZmaFOJCYmVkhDiRmZlaIA4mZmRXSaCCR9FpJMeLrTRP2WyLpM5J2S/qNpC9Ium9d4zYzs30OanoAfWuAvQN///GE538JeATwBuAu4MPAJuCpVQzOzMzGSyWQfDsibsryRElHA88GjomIf+hv2wVcJulZEfH1CsdpZmZD2lgjORb4l/kgAhARlwM7+4+ZmVmNUpmRXN2vcVwNnBYRf7HAc48EdozY/oP+Ywu63/3uF6tWrco1SDOzWbV9+/YbImL5qMeaDiQ/B94LXA4sAl4OnCHpkIg4fcw+y4A9I7bvBh466Q1XrVrFtm3b8o3WzGxGSbpm3GONBpKI2AJsGdi0WdK9gJMkfTQi7irjfSStA9YBrFy5soyXNDOzvhRrJF8G7gOsGvP4bmDJiO3L+o8dICI2RsRcRMwtXz5yZmZmZjmlGEhi6PuwHYyuhYyrnZiZWYVSDCQvBW4AxuXjNgMPkPSU+Q2S5ujVRzZXPzwzMxvUaI1E0tn0Cu3fpVdsP77/dcJ8fUTSj4CLIuL1ABHxLUlfAz4n6e3sW5B4sdeQmJnVr+murauA1wEPBgRcCbw6Ij4/8JyD6AWZQccDpwOfpjer+gpwQuWjtVbbdMUuTtlyFdft2cvhSxdz4nMeydqjVjQ9LLPW06zdandubi7c/lutFA/Ym67YxbvO+R57b7/z7m2LD17Eh1786MbHZtYGkrZHxNyox1KskViLzR+wd+3ZSwC79uzlXed8j01X7Gp0XKdsuWq/IAKw9/Y7OWXLVQ2NyKw7mk5tWccsdMCe9sy/zJnNdXv2TtxexUwqxdmZWdkcSCyXcQfILAfsrK8/mIqan9kAuQ7Ehy9dzK4RYzh86eJK3q+q1zRLkVNbNrWF0lfzB+Zh47aPU3Yq6sTnPJLFB+/fs7H44EWc+JxHVvJ+Vb2mWYocSGxqCx0gJx2wsyprZjNv7VEr+NCLH82KpYsRsGLp4v0K7WW/X1WvOc6mK3axesNWjlh/Pqs3bG28JmWzxamtGVB2nn6hA+T86xZ9v0mpqDzWHrVi7DiqeL8qXnMUp9CsaZ6RdFwVXVST0ldrj1rBJevXsHPD87lk/ZpcB7OyZjZNvl9dn8EpNGuaA0nHVXGQqeMAOSkVNY0saZ8y36/K1xylzhSa2ShObbVInhRVFQeZstJXsPBnWigVNc3rZ037lPF+w6p4zWF1pdDMxnEgaYm8efCqDjJ1H+TzKnNdS6pOfM4jR67aryoNaDbMqa2WyJuiqrvWMI06cvttT/s0lZab5v3NPCNpibwHxDxpqLK6vCa9Th0H+bwzshRWpDedlnM3mGXlQNISRVJU0xxkyjp4ZHmdOnL7edI+qRxAm07LNf3+1h5ObbVE21pJs7xOqt1fVabcpkkVNZ2Wa/r9rT08I2mJMjulFlLWwSPL69T1maZN+1R1AJ12ptN0N1bT72/t4UDSIlkOiEVz+2UdPLK+Th3tsdOq6gA6baqo6W6spt/f2sOprQ4pYxV7WemmlLvFJqlq7NPOdOpa0DhO0+9v7eEZSYeUURwtK9007euk0CU1r6qUW56ZTtMztqLvn9K/q1XHt9rtkCPWn8+of00BOzc8v+7hZDYrt8Gdlc85b9Y+b9f5Vrszoqx7gdRtVi46OGupoln5dzWntjqlrcXRWWozbTpVlVUZKalZ+neddQ4kHVJXO23ZUm0zndX8flkLMlP9d7XyNRZIJB0H/EfgCcAS4CrgzyPiryfsN6oMcFlEPLn8UbZPW854B6U4k0pldXsTylrRnvffdVYDeJs1OSN5K7ATeAtwA/A84IuS7hcRH5+w76nAlwf+/ttqhtgubf0FrGomVeTnMcuXBykrJZX3Om+zGsDbrMlA8h8i4oaBv2+VdDi9ADMpkPwkIi6tbmjt0/ZfwLJnUkV/HrOc3y8zJTXtv+ssB/A2a6xrayiIzLsCOLzusXSBO2T2V/Tn0dYOuDI0uZh0lgN4m6XW/ns08MMMz3u/pDsk3SDp05LuU/XAUudfwP0V/XmUeTBt2z09mmxTnuUA3mbJdG1JeiawFnjdhKf+JXAecD0wB7wXeKykJ0bEnaN2kLQOWAewcuXKsoacFHfI7K/oz6Osuk1bU45NNW0848jlfOHSa/dbWNt044VNlsTKdkmrgMuAf4yIF02577HA3wEviohNk56f2sr2Mm8ilfoq4jqbAVL5eazesHVkQFuxdDGXrF9T2zjaYNS/mYBXPnklf7b20c0NzICFV7Y3PiPpp6U2A9cAr8zxEl8FbgIeD2wqb2TVK/NsNfU1JHWfmafy83DKMbtRda0AvrHj+mYGZJk1GkgkHQJ8Bfg94I8i4pZpXyMiQhIw8jJTSSu7QyXlNSRNdOOk8POoIuXY1jbvSRx026uxYrukg4C/Af4AeG5E/DLn6zwX+H1ge4nDq8Us/eLM0mcdVHYHVBm3CiiqquYBF9rbq8murf9FbxHiB4H7SnrywNc9ASRdIOmC+R0krZO0UdLLJK2R9HbgLOBy4PwmPkQRs/SLM0ufdVDZHVBNt3lXGcjafA+bWddkauvZ/e8fHfHYEcBPgEVD268GXgO8BDgM+AXwOeC94zq2UpbipUGqMkufdViZKbamZ3ZVpihTqWvZ9BoLJBGxKsNznj709wuAC0Y/u31m6Rdnlj5rlZpu8646kKVQ17LpNd61NetS+cWpo4Cbymdts6Zndk0HMktTaivbrQEpFHAnadvq8Ko0fXMs1zFsFM9ILPkL5bV1dXhVmpzZ5b2ir1Oa3eZAYo0XcCdJPdClJLUUZRUnAQ5M6XFqy5JvzU090KUixRRl2e3KKX5GcyAx0s97px7oUtH0GpNRyj4JSPEzmgOJ0XwBd5LUA10qUpy5lX0SkOJnNNdIrC/l1twy16B0Ob+eYmtu2e3KKX5GcyCxligj0HW9+6vpNSajlL0QNcXPaA4kNqDLZ+vQ/e6vVK8eUOZsN9XPOOscSAzo/tk6zEZ+PeUUZVlm4TO2jYvtBsxGN0yZhV+vtDfbx4HEgNk4Wy+r+6vtaxkcBK1sTm01IMVaxCx0w5SVX29zraWsFGaK/4etOQ4kNUu1FjEr3TBl5NfbPHsrIwim+n/YmuPUVs1SrUWkvigxJSmvtJ+UtiojCKb6f9ia4xlJzVI+m3U3TDapzt6yzBTKSGGm/H/YmuEZSYmyFDFTPpu1bFKdvWWZKZTRcOD/wzbMM5KSZM0bp3o2a9NJcfaWZaZQRsOB/w/bMAeSkmQtYk76RU61GybVcdk+WdNWRYOgV5fbMAeSkkyTNx73i5xqN0yq47L91TlTSHFGZs1ptEYi6VGSLpB0i6TrJH1A0qIM+y2R9BlJuyX9RtIXJN23jjGPU0beONVumFTHZftLtXbTBC+6rFdjMxJJy4CvA1cCLwQeBpxKL7idNGH3LwGPAN4A3AV8GNgEPLWi4U5Uxtlgqt0wqY7LDuSZgmfQTWhyRvImYDHw4oj4+4g4AzgZeKukw8btJOlo4NnAayLi7Ig4F3gV8BRJz6pj4KOUcTaYajdMquMyG8Uz6Po1WSM5FtgSETcObDuL3uziGOC8Bfb7l4j4h/kNEXG5pJ39x75e0XgnKno2OM2sps7it7t0rE08g65fk4HkSGDr4IaIuFbSLf3HxgWSI4EdI7b/oP9YJeo4cGfthql76u4uHWuTWbhuXGqaDCTLgD0jtu/uP5Znv4cWHtUIdR64s8xqmrhoYB25d7cYWxk8g67fTLT/SloHrANYuXLl1PtPyrnWffDr4tTdBVIbJc/JhWfQ9WsykOwGlozYvqz/2EL7LZ9mv4jYCGwEmJubi+mGOf4APX+wq/vg18Wpe5svzW7VKHJy4e61ejXZtbWDoZqGpAcDhzC6BjJ2v75xtZPCxh2gF0mNdIeUdYOmlHRxlmXFuPuqPZoMJJuB50i698C244G9wEUT9nuApKfMb5A0R68+srmKgY47cN8Zoyc3VR/8urjwzC3GNswnF+3RZGrrDOAE4BxJH6YXCN4PnDbYEizpR8BFEfF6gIj4lqSvAZ+T9Hb2LUi8OCIqaf0dl3M9ZctVpaWYps0Fd23q7gKpDetiCrerGgskEbFb0jOBT9Br9d0DnE4vmAw6CBi+bMrx/ed+mt6s6iv0glJlxh24yzj4udDsAqkdyCcX7aEYk57pqrm5udi2bVtpr1dGy+rqDVtHnnmtWLqYS9avKWuoZq3jlvB0SNoeEXOjHpuJ9t+iFvrPPOv3ADerUtdSuF3lQDJBnrTTtGdRzgWbWZv5VrsTTNuCOB94du3ZS7Av8Cx0GesutvOa2ezwjGSCadNOeRbWudCcnXPmZulxIJlg2rRT3nqHc8GTubvNLE1ObU0wbdpp2oV1vpNbdl7pbJYmB5IJpl1FPk3gyVNPmWXubjNLk1NbGUyTdpqm3uELFU7H3W1maXIgqUDWwOMz7Ol4pbNZmpzaapAvVDidLl6s0qwLPCMpSZ62VJ9hT8/dbZaXW8er40BSgrxtqV4/YlYPt45Xy4GkBEWK5j7DNqueG1uq5UBSglSL5p7Km/Wk+jvaFS62lyDFornXqJjtk+LvaJc4kJQgxYsuehW42T4p/o52iVNbJUixaO6pvNk+Kf6OdokLine truncated
"text/plain": [
"<Figure size 432x288 with 1 Axes>"
]
},
"metadata": {
"needs_background": "light"
},
"output_type": "display_data"
}
],
"source": [
"plt.scatter(X_embedded[:,0], X_embedded[:,1], alpha=1.0)"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "50a0e3d3",
"metadata": {},
"outputs": [],
"source": []
}
],
"metadata": {
"kernelspec": {
"display_name": "Python 3 (ipykernel)",
"language": "python",
"name": "python3"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.10.6"
},
"vscode": {
"interpreter": {
"hash": "40e4fedf32e21c0b32c442446148826532e66f2735982b93501b453fc06f0092"
}
}
},
"nbformat": 4,
"nbformat_minor": 5
}