From b6f322dfa7f3f039d3c41ecb3f44e53f66469d95 Mon Sep 17 00:00:00 2001 From: gorold Date: Tue, 20 Sep 2022 17:01:39 +0800 Subject: [PATCH] fix bug in encoder/exponential smoothing module --- models/etsformer/encoder.py | 2 +- models/etsformer/exponential_smoothing.py | 4 ++-- models/etsformer/model.py | 2 +- 3 files changed, 4 insertions(+), 4 deletions(-) diff --git a/models/etsformer/encoder.py b/models/etsformer/encoder.py index 78a7619..d3f7173 100644 --- a/models/etsformer/encoder.py +++ b/models/etsformer/encoder.py @@ -35,7 +35,7 @@ class GrowthLayer(nn.Module): values = torch.cat([repeat(self.z0, 'h d -> b 1 h d', b=b), values], dim=1) values = values[:, 1:] - values[:, :-1] out = self.es(values) - out = torch.cat([repeat(self.es.v0, 'h d -> b 1 h d', b=b), out], dim=1) + out = torch.cat([repeat(self.es.v0, '1 1 h d -> b 1 h d', b=b), out], dim=1) out = rearrange(out, 'b t h d -> b t (h d)') return self.out_proj(out) diff --git a/models/etsformer/exponential_smoothing.py b/models/etsformer/exponential_smoothing.py index db8e82b..96c167d 100644 --- a/models/etsformer/exponential_smoothing.py +++ b/models/etsformer/exponential_smoothing.py @@ -60,8 +60,8 @@ class ExponentialSmoothing(nn.Module): # \alpha^t for all t = 1, 2, ..., T init_weight = self.weight ** (powers + 1) - return rearrange(init_weight, 'h t -> () t h ()'), \ - rearrange(weight, 'h t -> () t h ()') + return rearrange(init_weight, 'h t -> 1 t h 1'), \ + rearrange(weight, 'h t -> 1 t h 1') @property def weight(self): diff --git a/models/etsformer/model.py b/models/etsformer/model.py index fb0d028..89d0b2b 100644 --- a/models/etsformer/model.py +++ b/models/etsformer/model.py @@ -34,7 +34,7 @@ class ETSformer(nn.Module): self.configs = configs - assert configs.d_layers == configs.e_layers + assert configs.e_layers == configs.d_layers, "Encoder and decoder layers must be equal" # Embedding self.enc_embedding = ETSEmbedding(configs.enc_in, configs.d_model, dropout=configs.dropout)