mirror of
https://github.com/wassname/pytorch-transformer-ts.git
synced 2026-10-04 13:00:43 +08:00
initial torchscale models
This commit is contained in:
1 parent
efc4d0d678
commit
f6e7d77fe4
6 files changed
+1679
No files matched your search
@@ -11,3 +11,5 @@ einops
|
||||
opt_einsum
|
||||
pykeops
|
||||
scipy
|
||||
apex
|
||||
torchscale
|
||||
@@ -0,0 +1,10 @@
|
||||
# +
|
||||
from .estimator import TorchscaleEstimator
|
||||
from .lightning_module import TorchscaleightningModule
|
||||
from .module import TorchscaleModel
|
||||
|
||||
__all__ = [
|
||||
"TorchscaleModel",
|
||||
"TorchscaleightningModule",
|
||||
"TorchscaleEstimator",
|
||||
]
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
}
|
||||
Reference in new issue
Block a user