fixes
This commit is contained in:
Vendored
+201
@@ -0,0 +1,201 @@
|
||||
Apache License
|
||||
Version 2.0, January 2004
|
||||
http://www.apache.org/licenses/
|
||||
|
||||
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
||||
|
||||
1. Definitions.
|
||||
|
||||
"License" shall mean the terms and conditions for use, reproduction,
|
||||
and distribution as defined by Sections 1 through 9 of this document.
|
||||
|
||||
"Licensor" shall mean the copyright owner or entity authorized by
|
||||
the copyright owner that is granting the License.
|
||||
|
||||
"Legal Entity" shall mean the union of the acting entity and all
|
||||
other entities that control, are controlled by, or are under common
|
||||
control with that entity. For the purposes of this definition,
|
||||
"control" means (i) the power, direct or indirect, to cause the
|
||||
direction or management of such entity, whether by contract or
|
||||
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
||||
outstanding shares, or (iii) beneficial ownership of such entity.
|
||||
|
||||
"You" (or "Your") shall mean an individual or Legal Entity
|
||||
exercising permissions granted by this License.
|
||||
|
||||
"Source" form shall mean the preferred form for making modifications,
|
||||
including but not limited to software source code, documentation
|
||||
source, and configuration files.
|
||||
|
||||
"Object" form shall mean any form resulting from mechanical
|
||||
transformation or translation of a Source form, including but
|
||||
not limited to compiled object code, generated documentation,
|
||||
and conversions to other media types.
|
||||
|
||||
"Work" shall mean the work of authorship, whether in Source or
|
||||
Object form, made available under the License, as indicated by a
|
||||
copyright notice that is included in or attached to the work
|
||||
(an example is provided in the Appendix below).
|
||||
|
||||
"Derivative Works" shall mean any work, whether in Source or Object
|
||||
form, that is based on (or derived from) the Work and for which the
|
||||
editorial revisions, annotations, elaborations, or other modifications
|
||||
represent, as a whole, an original work of authorship. For the purposes
|
||||
of this License, Derivative Works shall not include works that remain
|
||||
separable from, or merely link (or bind by name) to the interfaces of,
|
||||
the Work and Derivative Works thereof.
|
||||
|
||||
"Contribution" shall mean any work of authorship, including
|
||||
the original version of the Work and any modifications or additions
|
||||
to that Work or Derivative Works thereof, that is intentionally
|
||||
submitted to Licensor for inclusion in the Work by the copyright owner
|
||||
or by an individual or Legal Entity authorized to submit on behalf of
|
||||
the copyright owner. For the purposes of this definition, "submitted"
|
||||
means any form of electronic, verbal, or written communication sent
|
||||
to the Licensor or its representatives, including but not limited to
|
||||
communication on electronic mailing lists, source code control systems,
|
||||
and issue tracking systems that are managed by, or on behalf of, the
|
||||
Licensor for the purpose of discussing and improving the Work, but
|
||||
excluding communication that is conspicuously marked or otherwise
|
||||
designated in writing by the copyright owner as "Not a Contribution."
|
||||
|
||||
"Contributor" shall mean Licensor and any individual or Legal Entity
|
||||
on behalf of whom a Contribution has been received by Licensor and
|
||||
subsequently incorporated within the Work.
|
||||
|
||||
2. Grant of Copyright License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
copyright license to reproduce, prepare Derivative Works of,
|
||||
publicly display, publicly perform, sublicense, and distribute the
|
||||
Work and such Derivative Works in Source or Object form.
|
||||
|
||||
3. Grant of Patent License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
(except as stated in this section) patent license to make, have made,
|
||||
use, offer to sell, sell, import, and otherwise transfer the Work,
|
||||
where such license applies only to those patent claims licensable
|
||||
by such Contributor that are necessarily infringed by their
|
||||
Contribution(s) alone or by combination of their Contribution(s)
|
||||
with the Work to which such Contribution(s) was submitted. If You
|
||||
institute patent litigation against any entity (including a
|
||||
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
||||
or a Contribution incorporated within the Work constitutes direct
|
||||
or contributory patent infringement, then any patent licenses
|
||||
granted to You under this License for that Work shall terminate
|
||||
as of the date such litigation is filed.
|
||||
|
||||
4. Redistribution. You may reproduce and distribute copies of the
|
||||
Work or Derivative Works thereof in any medium, with or without
|
||||
modifications, and in Source or Object form, provided that You
|
||||
meet the following conditions:
|
||||
|
||||
(a) You must give any other recipients of the Work or
|
||||
Derivative Works a copy of this License; and
|
||||
|
||||
(b) You must cause any modified files to carry prominent notices
|
||||
stating that You changed the files; and
|
||||
|
||||
(c) You must retain, in the Source form of any Derivative Works
|
||||
that You distribute, all copyright, patent, trademark, and
|
||||
attribution notices from the Source form of the Work,
|
||||
excluding those notices that do not pertain to any part of
|
||||
the Derivative Works; and
|
||||
|
||||
(d) If the Work includes a "NOTICE" text file as part of its
|
||||
distribution, then any Derivative Works that You distribute must
|
||||
include a readable copy of the attribution notices contained
|
||||
within such NOTICE file, excluding those notices that do not
|
||||
pertain to any part of the Derivative Works, in at least one
|
||||
of the following places: within a NOTICE text file distributed
|
||||
as part of the Derivative Works; within the Source form or
|
||||
documentation, if provided along with the Derivative Works; or,
|
||||
within a display generated by the Derivative Works, if and
|
||||
wherever such third-party notices normally appear. The contents
|
||||
of the NOTICE file are for informational purposes only and
|
||||
do not modify the License. You may add Your own attribution
|
||||
notices within Derivative Works that You distribute, alongside
|
||||
or as an addendum to the NOTICE text from the Work, provided
|
||||
that such additional attribution notices cannot be construed
|
||||
as modifying the License.
|
||||
|
||||
You may add Your own copyright statement to Your modifications and
|
||||
may provide additional or different license terms and conditions
|
||||
for use, reproduction, or distribution of Your modifications, or
|
||||
for any such Derivative Works as a whole, provided Your use,
|
||||
reproduction, and distribution of the Work otherwise complies with
|
||||
the conditions stated in this License.
|
||||
|
||||
5. Submission of Contributions. Unless You explicitly state otherwise,
|
||||
any Contribution intentionally submitted for inclusion in the Work
|
||||
by You to the Licensor shall be under the terms and conditions of
|
||||
this License, without any additional terms or conditions.
|
||||
Notwithstanding the above, nothing herein shall supersede or modify
|
||||
the terms of any separate license agreement you may have executed
|
||||
with Licensor regarding such Contributions.
|
||||
|
||||
6. Trademarks. This License does not grant permission to use the trade
|
||||
names, trademarks, service marks, or product names of the Licensor,
|
||||
except as required for reasonable and customary use in describing the
|
||||
origin of the Work and reproducing the content of the NOTICE file.
|
||||
|
||||
7. Disclaimer of Warranty. Unless required by applicable law or
|
||||
agreed to in writing, Licensor provides the Work (and each
|
||||
Contributor provides its Contributions) on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
||||
implied, including, without limitation, any warranties or conditions
|
||||
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
||||
PARTICULAR PURPOSE. You are solely responsible for determining the
|
||||
appropriateness of using or redistributing the Work and assume any
|
||||
risks associated with Your exercise of permissions under this License.
|
||||
|
||||
8. Limitation of Liability. In no event and under no legal theory,
|
||||
whether in tort (including negligence), contract, or otherwise,
|
||||
unless required by applicable law (such as deliberate and grossly
|
||||
negligent acts) or agreed to in writing, shall any Contributor be
|
||||
liable to You for damages, including any direct, indirect, special,
|
||||
incidental, or consequential damages of any character arising as a
|
||||
result of this License or out of the use or inability to use the
|
||||
Work (including but not limited to damages for loss of goodwill,
|
||||
work stoppage, computer failure or malfunction, or any and all
|
||||
other commercial damages or losses), even if such Contributor
|
||||
has been advised of the possibility of such damages.
|
||||
|
||||
9. Accepting Warranty or Additional Liability. While redistributing
|
||||
the Work or Derivative Works thereof, You may choose to offer,
|
||||
and charge a fee for, acceptance of support, warranty, indemnity,
|
||||
or other liability obligations and/or rights consistent with this
|
||||
License. However, in accepting such obligations, You may act only
|
||||
on Your own behalf and on Your sole responsibility, not on behalf
|
||||
of any other Contributor, and only if You agree to indemnify,
|
||||
defend, and hold each Contributor harmless for any liability
|
||||
incurred by, or claims asserted against, such Contributor by reason
|
||||
of your accepting any such warranty or additional liability.
|
||||
|
||||
END OF TERMS AND CONDITIONS
|
||||
|
||||
APPENDIX: How to apply the Apache License to your work.
|
||||
|
||||
To apply the Apache License to your work, attach the following
|
||||
boilerplate notice, with the fields enclosed by brackets "[]"
|
||||
replaced with your own identifying information. (Don't include
|
||||
the brackets!) The text should be enclosed in the appropriate
|
||||
comment syntax for the file format. We also recommend that a
|
||||
file or class name and description of purpose be included on the
|
||||
same "printed page" as the copyright notice for easier
|
||||
identification within third-party archives.
|
||||
|
||||
Copyright 2021-2022 NVIDIA Corporation
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
Vendored
+5
@@ -0,0 +1,5 @@
|
||||
TFT for PyTorch
|
||||
|
||||
This repository includes software from https://github.com/google-research/google-research/tree/master/tft licensed under the Apache 2.0 License.
|
||||
|
||||
This repository contains code from https://github.com/rwightman/pytorch-image-models/blob/master/timm/utils/model_ema.py under the Apache 2.0 License.
|
||||
+3
@@ -0,0 +1,3 @@
|
||||
This folder contains code copied from NVIDIA's Temporal Fusion Transformer implementation, licensed under Apache 2.0.
|
||||
All rights belong to NVIDIA Corporation.
|
||||
Modifications are noted in the file headers.
|
||||
+525
@@ -0,0 +1,525 @@
|
||||
# Copyright (c) 2021-2022, NVIDIA CORPORATION. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
# Modified by Alexander Blank, 2025.
|
||||
# Modifications:
|
||||
# - added support for multiple outputs
|
||||
# - added support for mode configurable targets
|
||||
# - added support for single dimension, non-quantile outputs
|
||||
# - added support for target agnostic predictions, for cases, where the target does not become known after prediction
|
||||
|
||||
|
||||
import os
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from torch import Tensor
|
||||
from torch.nn.parameter import UninitializedParameter
|
||||
from typing import Dict, Tuple, Optional, List
|
||||
|
||||
|
||||
MAKE_CONVERT_COMPATIBLE = os.environ.get("TFT_SCRIPTING", None) is not None
|
||||
from torch.nn import LayerNorm
|
||||
|
||||
|
||||
class MaybeLayerNorm(nn.Module):
|
||||
def __init__(self, output_size, hidden_size, eps):
|
||||
super().__init__()
|
||||
if output_size and output_size == 1:
|
||||
self.ln = nn.Identity()
|
||||
else:
|
||||
self.ln = LayerNorm(output_size if output_size else hidden_size, eps=eps)
|
||||
|
||||
def forward(self, x):
|
||||
return self.ln(x)
|
||||
|
||||
|
||||
class GLU(nn.Module):
|
||||
def __init__(self, hidden_size, output_size):
|
||||
super().__init__()
|
||||
self.lin = nn.Linear(hidden_size, output_size * 2)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
x = self.lin(x)
|
||||
x = F.glu(x)
|
||||
return x
|
||||
|
||||
|
||||
class GRN(nn.Module):
|
||||
def __init__(self,
|
||||
input_size,
|
||||
hidden_size,
|
||||
output_size=None,
|
||||
context_hidden_size=None,
|
||||
dropout=0.0, ):
|
||||
super().__init__()
|
||||
self.layer_norm = MaybeLayerNorm(output_size, hidden_size, eps=1e-3)
|
||||
self.lin_a = nn.Linear(input_size, hidden_size)
|
||||
if context_hidden_size is not None:
|
||||
self.lin_c = nn.Linear(context_hidden_size, hidden_size, bias=False)
|
||||
else:
|
||||
self.lin_c = nn.Identity()
|
||||
self.lin_i = nn.Linear(hidden_size, hidden_size)
|
||||
self.glu = GLU(hidden_size, output_size if output_size else hidden_size)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
self.out_proj = nn.Linear(input_size, output_size) if output_size else None
|
||||
|
||||
def forward(self, a: Tensor, c: Optional[Tensor] = None):
|
||||
x = self.lin_a(a)
|
||||
if c is not None:
|
||||
x = x + self.lin_c(c).unsqueeze(1)
|
||||
x = F.elu(x)
|
||||
x = self.lin_i(x)
|
||||
x = self.dropout(x)
|
||||
x = self.glu(x)
|
||||
y = a if self.out_proj is None else self.out_proj(a)
|
||||
x = x + y
|
||||
return self.layer_norm(x)
|
||||
|
||||
# @torch.jit.script #Currently broken with autocast
|
||||
|
||||
|
||||
def fused_pointwise_linear_v1(x, a, b):
|
||||
out = torch.mul(x.unsqueeze(-1), a)
|
||||
out = out + b
|
||||
return out
|
||||
|
||||
|
||||
@torch.jit.script
|
||||
def fused_pointwise_linear_v2(x, a, b):
|
||||
out = x.unsqueeze(3) * a
|
||||
out = out + b
|
||||
return out
|
||||
|
||||
|
||||
class TFTEmbedding(nn.Module):
|
||||
def __init__(self, config, initialize_cont_params=True):
|
||||
# initialize_cont_params=False prevents form initializing parameters inside this class
|
||||
# so they can be lazily initialized in LazyEmbedding module
|
||||
super().__init__()
|
||||
self.s_cat_inp_lens = config.static_categorical_inp_lens
|
||||
self.t_cat_k_inp_lens = config.temporal_known_categorical_inp_lens
|
||||
self.t_cat_o_inp_lens = config.temporal_observed_categorical_inp_lens
|
||||
self.s_cont_inp_size = config.static_continuous_inp_size
|
||||
self.t_cont_k_inp_size = config.temporal_known_continuous_inp_size
|
||||
self.t_cont_o_inp_size = config.temporal_observed_continuous_inp_size
|
||||
self.t_tgt_size = config.temporal_target_size
|
||||
|
||||
self.hidden_size = config.hidden_size
|
||||
|
||||
# There are 7 types of input:
|
||||
# 1. Static categorical
|
||||
# 2. Static continuous
|
||||
# 3. Temporal known a priori categorical
|
||||
# 4. Temporal known a priori continuous
|
||||
# 5. Temporal observed categorical
|
||||
# 6. Temporal observed continuous
|
||||
# 7. Temporal observed targets (time series obseved so far)
|
||||
|
||||
self.s_cat_embed = nn.ModuleList([
|
||||
nn.Embedding(n, self.hidden_size) for n in self.s_cat_inp_lens]) if self.s_cat_inp_lens else None
|
||||
self.t_cat_k_embed = nn.ModuleList([
|
||||
nn.Embedding(n, self.hidden_size) for n in self.t_cat_k_inp_lens]) if self.t_cat_k_inp_lens else None
|
||||
self.t_cat_o_embed = nn.ModuleList([
|
||||
nn.Embedding(n, self.hidden_size) for n in self.t_cat_o_inp_lens]) if self.t_cat_o_inp_lens else None
|
||||
|
||||
if initialize_cont_params:
|
||||
self.s_cont_embedding_vectors = nn.Parameter(
|
||||
torch.Tensor(self.s_cont_inp_size, self.hidden_size)) if self.s_cont_inp_size else None
|
||||
self.t_cont_k_embedding_vectors = nn.Parameter(
|
||||
torch.Tensor(self.t_cont_k_inp_size, self.hidden_size)) if self.t_cont_k_inp_size else None
|
||||
self.t_cont_o_embedding_vectors = nn.Parameter(
|
||||
torch.Tensor(self.t_cont_o_inp_size, self.hidden_size)) if self.t_cont_o_inp_size else None
|
||||
self.t_tgt_embedding_vectors = nn.Parameter(torch.Tensor(self.t_tgt_size, self.hidden_size))
|
||||
|
||||
self.s_cont_embedding_bias = nn.Parameter(
|
||||
torch.zeros(self.s_cont_inp_size, self.hidden_size)) if self.s_cont_inp_size else None
|
||||
self.t_cont_k_embedding_bias = nn.Parameter(
|
||||
torch.zeros(self.t_cont_k_inp_size, self.hidden_size)) if self.t_cont_k_inp_size else None
|
||||
self.t_cont_o_embedding_bias = nn.Parameter(
|
||||
torch.zeros(self.t_cont_o_inp_size, self.hidden_size)) if self.t_cont_o_inp_size else None
|
||||
self.t_tgt_embedding_bias = nn.Parameter(torch.zeros(self.t_tgt_size, self.hidden_size))
|
||||
|
||||
self.reset_parameters()
|
||||
|
||||
def reset_parameters(self):
|
||||
if self.s_cont_embedding_vectors is not None:
|
||||
torch.nn.init.xavier_normal_(self.s_cont_embedding_vectors)
|
||||
torch.nn.init.zeros_(self.s_cont_embedding_bias)
|
||||
if self.t_cont_k_embedding_vectors is not None:
|
||||
torch.nn.init.xavier_normal_(self.t_cont_k_embedding_vectors)
|
||||
torch.nn.init.zeros_(self.t_cont_k_embedding_bias)
|
||||
if self.t_cont_o_embedding_vectors is not None:
|
||||
torch.nn.init.xavier_normal_(self.t_cont_o_embedding_vectors)
|
||||
torch.nn.init.zeros_(self.t_cont_o_embedding_bias)
|
||||
if self.t_tgt_embedding_vectors is not None:
|
||||
torch.nn.init.xavier_normal_(self.t_tgt_embedding_vectors)
|
||||
torch.nn.init.zeros_(self.t_tgt_embedding_bias)
|
||||
if self.s_cat_embed is not None:
|
||||
for module in self.s_cat_embed:
|
||||
module.reset_parameters()
|
||||
if self.t_cat_k_embed is not None:
|
||||
for module in self.t_cat_k_embed:
|
||||
module.reset_parameters()
|
||||
if self.t_cat_o_embed is not None:
|
||||
for module in self.t_cat_o_embed:
|
||||
module.reset_parameters()
|
||||
|
||||
def _apply_embedding(self,
|
||||
cat: Optional[Tensor],
|
||||
cont: Optional[Tensor],
|
||||
cat_emb: Optional[nn.ModuleList],
|
||||
cont_emb: Tensor,
|
||||
cont_bias: Tensor,
|
||||
) -> Tuple[Optional[Tensor], Optional[Tensor]]:
|
||||
e_cat = torch.stack([embed(cat[..., i]) for i, embed in enumerate(cat_emb)],
|
||||
dim=-2) if cat is not None else None
|
||||
if cont is not None:
|
||||
# the line below is equivalent to following einsums
|
||||
# e_cont = torch.einsum('btf,fh->bthf', cont, cont_emb)
|
||||
# e_cont = torch.einsum('bf,fh->bhf', cont, cont_emb)
|
||||
if MAKE_CONVERT_COMPATIBLE:
|
||||
e_cont = torch.mul(cont.unsqueeze(-1), cont_emb)
|
||||
e_cont = e_cont + cont_bias
|
||||
else:
|
||||
e_cont = fused_pointwise_linear_v1(cont, cont_emb, cont_bias)
|
||||
else:
|
||||
e_cont = None
|
||||
|
||||
if e_cat is not None and e_cont is not None:
|
||||
return torch.cat([e_cat, e_cont], dim=-2)
|
||||
elif e_cat is not None:
|
||||
return e_cat
|
||||
elif e_cont is not None:
|
||||
return e_cont
|
||||
else:
|
||||
return None
|
||||
|
||||
def forward(self, x: Dict[str, Tensor], use_target: bool = False):
|
||||
# Extract inputs
|
||||
s_cat_inp = x.get('s_cat', None)
|
||||
s_cont_inp = x.get('s_cont', None)
|
||||
t_cat_k_inp = x.get('k_cat', None)
|
||||
t_cont_k_inp = x.get('k_cont', None)
|
||||
t_cat_o_inp = x.get('o_cat', None)
|
||||
t_cont_o_inp = x.get('o_cont', None)
|
||||
|
||||
# Only use target if teacher forcing is enabled.
|
||||
# When disabled, we ignore target values.
|
||||
if use_target:
|
||||
t_tgt_obs = x['target'] # Must be present when using teacher forcing
|
||||
else:
|
||||
t_tgt_obs = None
|
||||
|
||||
# For static inputs, take the first timestep
|
||||
s_cat_inp = s_cat_inp[:, 0, :] if s_cat_inp is not None else None
|
||||
s_cont_inp = s_cont_inp[:, 0, :] if s_cont_inp is not None else None
|
||||
|
||||
# Apply embeddings for static and known/observed temporal features
|
||||
s_inp = self._apply_embedding(s_cat_inp,
|
||||
s_cont_inp,
|
||||
self.s_cat_embed,
|
||||
self.s_cont_embedding_vectors,
|
||||
self.s_cont_embedding_bias)
|
||||
t_known_inp = self._apply_embedding(t_cat_k_inp,
|
||||
t_cont_k_inp,
|
||||
self.t_cat_k_embed,
|
||||
self.t_cont_k_embedding_vectors,
|
||||
self.t_cont_k_embedding_bias)
|
||||
t_observed_inp = self._apply_embedding(t_cat_o_inp,
|
||||
t_cont_o_inp,
|
||||
self.t_cat_o_embed,
|
||||
self.t_cont_o_embedding_vectors,
|
||||
self.t_cont_o_embedding_bias)
|
||||
# Compute the target embedding only if teacher forcing is enabled.
|
||||
if use_target and t_tgt_obs is not None:
|
||||
if MAKE_CONVERT_COMPATIBLE:
|
||||
t_observed_tgt = torch.matmul(t_tgt_obs.unsqueeze(3).unsqueeze(4),
|
||||
self.t_tgt_embedding_vectors.unsqueeze(1)).squeeze(3)
|
||||
t_observed_tgt = t_observed_tgt + self.t_tgt_embedding_bias
|
||||
else:
|
||||
t_observed_tgt = fused_pointwise_linear_v2(t_tgt_obs,
|
||||
self.t_tgt_embedding_vectors,
|
||||
self.t_tgt_embedding_bias)
|
||||
else:
|
||||
t_observed_tgt = None
|
||||
|
||||
return s_inp, t_known_inp, t_observed_inp, t_observed_tgt
|
||||
|
||||
|
||||
class LazyEmbedding(nn.modules.lazy.LazyModuleMixin, TFTEmbedding):
|
||||
cls_to_become = TFTEmbedding
|
||||
|
||||
def __init__(self, config):
|
||||
super().__init__(config, initialize_cont_params=False)
|
||||
|
||||
if config.static_continuous_inp_size:
|
||||
self.s_cont_embedding_vectors = UninitializedParameter()
|
||||
self.s_cont_embedding_bias = UninitializedParameter()
|
||||
else:
|
||||
self.s_cont_embedding_vectors = None
|
||||
self.s_cont_embedding_bias = None
|
||||
|
||||
if config.temporal_known_continuous_inp_size:
|
||||
self.t_cont_k_embedding_vectors = UninitializedParameter()
|
||||
self.t_cont_k_embedding_bias = UninitializedParameter()
|
||||
else:
|
||||
self.t_cont_k_embedding_vectors = None
|
||||
self.t_cont_k_embedding_bias = None
|
||||
|
||||
if config.temporal_observed_continuous_inp_size:
|
||||
self.t_cont_o_embedding_vectors = UninitializedParameter()
|
||||
self.t_cont_o_embedding_bias = UninitializedParameter()
|
||||
else:
|
||||
self.t_cont_o_embedding_vectors = None
|
||||
self.t_cont_o_embedding_bias = None
|
||||
|
||||
self.t_tgt_embedding_vectors = UninitializedParameter()
|
||||
self.t_tgt_embedding_bias = UninitializedParameter()
|
||||
|
||||
def initialize_parameters(self, x):
|
||||
if self.has_uninitialized_params():
|
||||
s_cont_inp = x.get('s_cont', None)
|
||||
t_cont_k_inp = x.get('k_cont', None)
|
||||
t_cont_o_inp = x.get('o_cont', None)
|
||||
t_tgt_obs = x['target'] # Has to be present
|
||||
|
||||
if s_cont_inp is not None:
|
||||
self.s_cont_embedding_vectors.materialize((s_cont_inp.shape[-1], self.hidden_size))
|
||||
self.s_cont_embedding_bias.materialize((s_cont_inp.shape[-1], self.hidden_size))
|
||||
|
||||
if t_cont_k_inp is not None:
|
||||
self.t_cont_k_embedding_vectors.materialize((t_cont_k_inp.shape[-1], self.hidden_size))
|
||||
self.t_cont_k_embedding_bias.materialize((t_cont_k_inp.shape[-1], self.hidden_size))
|
||||
|
||||
if t_cont_o_inp is not None:
|
||||
self.t_cont_o_embedding_vectors.materialize((t_cont_o_inp.shape[-1], self.hidden_size))
|
||||
self.t_cont_o_embedding_bias.materialize((t_cont_o_inp.shape[-1], self.hidden_size))
|
||||
|
||||
self.t_tgt_embedding_vectors.materialize((t_tgt_obs.shape[-1], self.hidden_size))
|
||||
self.t_tgt_embedding_bias.materialize((t_tgt_obs.shape[-1], self.hidden_size))
|
||||
|
||||
self.reset_parameters()
|
||||
|
||||
# def forward(self, x: Dict[str, Tensor], use_target: bool = True):
|
||||
# return super().forward(x, use_target=use_target)
|
||||
|
||||
|
||||
class VariableSelectionNetwork(nn.Module):
|
||||
def __init__(self, config, num_inputs):
|
||||
super().__init__()
|
||||
self.joint_grn = GRN(config.hidden_size * num_inputs, config.hidden_size, output_size=num_inputs,
|
||||
context_hidden_size=config.hidden_size)
|
||||
self.var_grns = nn.ModuleList(
|
||||
[GRN(config.hidden_size, config.hidden_size, dropout=config.dropout) for _ in range(num_inputs)])
|
||||
|
||||
def forward(self, x: Tensor, context: Optional[Tensor] = None):
|
||||
Xi = torch.flatten(x, start_dim=-2)
|
||||
grn_outputs = self.joint_grn(Xi, c=context)
|
||||
sparse_weights = F.softmax(grn_outputs, dim=-1)
|
||||
transformed_embed_list = [m(x[..., i, :]) for i, m in enumerate(self.var_grns)]
|
||||
transformed_embed = torch.stack(transformed_embed_list, dim=-1)
|
||||
# the line below performs batched matrix vector multiplication
|
||||
# for temporal features it's bthf,btf->bth
|
||||
# for static features it's bhf,bf->bh
|
||||
variable_ctx = torch.matmul(transformed_embed, sparse_weights.unsqueeze(-1)).squeeze(-1)
|
||||
|
||||
return variable_ctx, sparse_weights
|
||||
|
||||
|
||||
class StaticCovariateEncoder(nn.Module):
|
||||
def __init__(self, config):
|
||||
super().__init__()
|
||||
self.vsn = VariableSelectionNetwork(config, config.num_static_vars)
|
||||
self.context_grns = nn.ModuleList(
|
||||
[GRN(config.hidden_size, config.hidden_size, dropout=config.dropout) for _ in range(4)])
|
||||
|
||||
def forward(self, x: Tensor) -> Tuple[Tensor, Tensor, Tensor, Tensor]:
|
||||
variable_ctx, sparse_weights = self.vsn(x)
|
||||
|
||||
# Context vectors:
|
||||
# variable selection context
|
||||
# enrichment context
|
||||
# state_c context
|
||||
# state_h context
|
||||
cs, ce, ch, cc = [m(variable_ctx) for m in self.context_grns]
|
||||
|
||||
return cs, ce, ch, cc
|
||||
|
||||
|
||||
class InterpretableMultiHeadAttention(nn.Module):
|
||||
def __init__(self, config):
|
||||
super().__init__()
|
||||
self.n_head = config.n_head
|
||||
assert config.hidden_size % config.n_head == 0
|
||||
self.d_head = config.hidden_size // config.n_head
|
||||
self.qkv_linears = nn.Linear(config.hidden_size, (2 * self.n_head + 1) * self.d_head, bias=False)
|
||||
self.out_proj = nn.Linear(self.d_head, config.hidden_size, bias=False)
|
||||
self.attn_dropout = nn.Dropout(config.attn_dropout)
|
||||
self.out_dropout = nn.Dropout(config.dropout)
|
||||
self.scale = self.d_head ** -0.5
|
||||
self.register_buffer("_mask",
|
||||
torch.triu(torch.full((config.example_length, config.example_length), float('-inf')),
|
||||
1).unsqueeze(0))
|
||||
|
||||
def forward(self, x: Tensor) -> Tuple[Tensor, Tensor]:
|
||||
bs, t, h_size = x.shape
|
||||
qkv = self.qkv_linears(x)
|
||||
q, k, v = qkv.split((self.n_head * self.d_head, self.n_head * self.d_head, self.d_head), dim=-1)
|
||||
q = q.view(bs, t, self.n_head, self.d_head)
|
||||
k = k.view(bs, t, self.n_head, self.d_head)
|
||||
v = v.view(bs, t, self.d_head)
|
||||
|
||||
# attn_score = torch.einsum('bind,bjnd->bnij', q, k)
|
||||
attn_score = torch.matmul(q.permute((0, 2, 1, 3)), k.permute((0, 2, 3, 1)))
|
||||
attn_score.mul_(self.scale)
|
||||
|
||||
attn_score = attn_score + self._mask
|
||||
|
||||
attn_prob = F.softmax(attn_score, dim=3)
|
||||
attn_prob = self.attn_dropout(attn_prob)
|
||||
|
||||
# attn_vec = torch.einsum('bnij,bjd->bnid', attn_prob, v)
|
||||
attn_vec = torch.matmul(attn_prob, v.unsqueeze(1))
|
||||
m_attn_vec = torch.mean(attn_vec, dim=1)
|
||||
out = self.out_proj(m_attn_vec)
|
||||
out = self.out_dropout(out)
|
||||
|
||||
return out, attn_prob
|
||||
|
||||
|
||||
class TFTBack(nn.Module):
|
||||
def __init__(self, config):
|
||||
super().__init__()
|
||||
|
||||
self.encoder_length = config.encoder_length
|
||||
self.history_vsn = VariableSelectionNetwork(config, config.num_historic_vars)
|
||||
self.history_encoder = nn.LSTM(config.hidden_size, config.hidden_size, batch_first=True)
|
||||
self.future_vsn = VariableSelectionNetwork(config, config.num_future_vars)
|
||||
self.future_encoder = nn.LSTM(config.hidden_size, config.hidden_size, batch_first=True)
|
||||
|
||||
self.input_gate = GLU(config.hidden_size, config.hidden_size)
|
||||
self.input_gate_ln = LayerNorm(config.hidden_size, eps=1e-3)
|
||||
|
||||
self.enrichment_grn = GRN(config.hidden_size,
|
||||
config.hidden_size,
|
||||
context_hidden_size=config.hidden_size,
|
||||
dropout=config.dropout)
|
||||
self.attention = InterpretableMultiHeadAttention(config)
|
||||
self.attention_gate = GLU(config.hidden_size, config.hidden_size)
|
||||
self.attention_ln = LayerNorm(config.hidden_size, eps=1e-3)
|
||||
|
||||
self.positionwise_grn = GRN(config.hidden_size,
|
||||
config.hidden_size,
|
||||
dropout=config.dropout)
|
||||
|
||||
self.decoder_gate = GLU(config.hidden_size, config.hidden_size)
|
||||
self.decoder_ln = LayerNorm(config.hidden_size, eps=1e-3)
|
||||
|
||||
self.quantiles = config.quantiles
|
||||
self.target_size = config.target_size
|
||||
if self.quantiles is not None:
|
||||
self.output = nn.Linear(config.hidden_size, len(config.quantiles) * config.target_size)
|
||||
else:
|
||||
self.output = nn.Linear(config.hidden_size, config.target_size)
|
||||
|
||||
def forward(self, historical_inputs, cs, ch, cc, ce, future_inputs):
|
||||
historical_features, _ = self.history_vsn(historical_inputs, cs)
|
||||
history, state = self.history_encoder(historical_features, (ch, cc))
|
||||
future_features, _ = self.future_vsn(future_inputs, cs)
|
||||
future, _ = self.future_encoder(future_features, state)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# skip connection
|
||||
input_embedding = torch.cat([historical_features, future_features], dim=1)
|
||||
temporal_features = torch.cat([history, future], dim=1)
|
||||
temporal_features = self.input_gate(temporal_features)
|
||||
temporal_features = temporal_features + input_embedding
|
||||
temporal_features = self.input_gate_ln(temporal_features)
|
||||
|
||||
# Static enrichment
|
||||
enriched = self.enrichment_grn(temporal_features, c=ce)
|
||||
|
||||
# Temporal self attention
|
||||
x, _ = self.attention(enriched)
|
||||
|
||||
# Don't compute hictorical quantiles
|
||||
x = x[:, self.encoder_length:, :]
|
||||
temporal_features = temporal_features[:, self.encoder_length:, :]
|
||||
enriched = enriched[:, self.encoder_length:, :]
|
||||
|
||||
x = self.attention_gate(x)
|
||||
x = x + enriched
|
||||
x = self.attention_ln(x)
|
||||
|
||||
# Position-wise feed-forward
|
||||
x = self.positionwise_grn(x)
|
||||
|
||||
# Final skip connection
|
||||
x = self.decoder_gate(x)
|
||||
x = x + temporal_features
|
||||
x = self.decoder_ln(x)
|
||||
|
||||
out = self.output(x)
|
||||
if self.quantiles is not None:
|
||||
# Reshape to [batch, time, target_size, n_quantiles]
|
||||
out = out.view(out.size(0), out.size(1), self.target_size, len(self.quantiles))
|
||||
else:
|
||||
# Reshape to [batch, time, target_size]
|
||||
out = out.view(out.size(0), out.size(1), self.target_size)
|
||||
|
||||
return out
|
||||
|
||||
|
||||
class TemporalFusionTransformer(nn.Module):
|
||||
"""
|
||||
Implementation of https://arxiv.org/abs/1912.09363
|
||||
"""
|
||||
|
||||
def __init__(self, config):
|
||||
super().__init__()
|
||||
|
||||
if hasattr(config, 'model'):
|
||||
config = config.model
|
||||
|
||||
self.encoder_length = config.encoder_length # this determines from how distant past we want to use data from
|
||||
|
||||
# self.embedding = LazyEmbedding(config)
|
||||
self.embedding = TFTEmbedding(config)
|
||||
self.static_encoder = StaticCovariateEncoder(config)
|
||||
# if MAKE_CONVERT_COMPATIBLE:
|
||||
self.TFTpart2 = TFTBack(config)
|
||||
# else:
|
||||
# self.TFTpart2 = torch.jit.script(TFTBack(config))
|
||||
|
||||
def forward(self, x: Dict[str, Tensor]) -> Tensor:
|
||||
# Call embedding with use_target=False to skip target features entirely.
|
||||
s_inp, t_known_inp, t_observed_inp, t_observed_tgt = self.embedding(x, use_target=False)
|
||||
|
||||
# Compute static context
|
||||
cs, ce, ch, cc = self.static_encoder(s_inp)
|
||||
ch, cc = ch.unsqueeze(0), cc.unsqueeze(0) # Initialize LSTM states
|
||||
|
||||
# Build historical inputs without teacher-forced targets.
|
||||
# Include observed features if available, and the known inputs.
|
||||
historical_inputs = []
|
||||
if t_observed_inp is not None:
|
||||
historical_inputs.append(t_observed_inp[:, :self.encoder_length, :])
|
||||
historical_inputs.append(t_known_inp[:, :self.encoder_length, :])
|
||||
historical_inputs = torch.cat(historical_inputs, dim=-2)
|
||||
|
||||
# Future inputs remain the same
|
||||
future_inputs = t_known_inp[:, self.encoder_length:]
|
||||
return self.TFTpart2(historical_inputs, cs, ch, cc, ce, future_inputs)
|
||||
Reference in New Issue
Block a user