Skip to content

Commit f294c00

Browse files
[ENH] ModernTCN Implemented
1 parent bd803bf commit f294c00

3 files changed

Lines changed: 267 additions & 40 deletions

File tree

Lines changed: 194 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -1,29 +1,207 @@
11
"""
2-
ModernTCN block: DWConv + ConvFFN + residual.
2+
ModernTCN building blocks: ReparamLargeKernelConv, ModernTCNBlock, and Flatten_Head.
33
"""
44

5+
import torch
56
import torch.nn as nn
67

78

89
class ModernTCNBlock(nn.Module):
9-
def __init__(self, d_model, kernel_size, d_ff, dropout):
10+
"""
11+
Modern TCN Block.
12+
13+
This block is a residual block that consists of a depthwise
14+
separable convolution and a feed-forward network.
15+
16+
Parameters
17+
----------
18+
d_model : int
19+
Dimension of the model.
20+
kernel_size : int
21+
Size of the large kernel.
22+
small_kernel_size : int
23+
Size of the small kernel.
24+
d_ff : int
25+
Dimension of the feed-forward network.
26+
nvars : int
27+
Number of variables.
28+
dropout : float
29+
Dropout rate.
30+
"""
31+
32+
def __init__(self, d_model, kernel_size, small_kernel_size, d_ff, nvars, dropout):
1033
super().__init__()
11-
self.dwconv = nn.Conv1d(
12-
d_model,
13-
d_model,
34+
self.d_model = d_model
35+
self.nvars = nvars
36+
37+
self.dwconv = ReparamLargeKernelConv(
38+
in_channels=nvars * d_model,
39+
out_channels=nvars * d_model,
1440
kernel_size=kernel_size,
15-
padding=kernel_size // 2,
16-
groups=d_model,
41+
stride=1,
42+
groups=nvars * d_model,
43+
small_kernel_size=small_kernel_size,
1744
)
1845
self.norm = nn.BatchNorm1d(d_model)
19-
self.pw_conv1 = nn.Conv1d(d_model, d_ff, kernel_size=1)
20-
self.pw_conv2 = nn.Conv1d(d_ff, d_model, kernel_size=1)
21-
self.act = nn.GELU()
22-
self.drop = nn.Dropout(dropout)
46+
47+
self.ffn1_pw1 = nn.Conv1d(
48+
nvars * d_model, nvars * d_ff, kernel_size=1, groups=nvars
49+
)
50+
self.ffn1_act = nn.GELU()
51+
self.ffn1_pw2 = nn.Conv1d(
52+
nvars * d_ff, nvars * d_model, kernel_size=1, groups=nvars
53+
)
54+
self.ffn1_drop1 = nn.Dropout(dropout)
55+
self.ffn1_drop2 = nn.Dropout(dropout)
56+
57+
self.ffn2_pw1 = nn.Conv1d(
58+
nvars * d_model, nvars * d_ff, kernel_size=1, groups=d_model
59+
)
60+
self.ffn2_act = nn.GELU()
61+
self.ffn2_pw2 = nn.Conv1d(
62+
nvars * d_ff, nvars * d_model, kernel_size=1, groups=d_model
63+
)
64+
self.ffn2_drop1 = nn.Dropout(dropout)
65+
self.ffn2_drop2 = nn.Dropout(dropout)
66+
67+
def forward(self, x):
68+
input_x = x
69+
B, M, D, N = x.shape
70+
71+
x = x.reshape(B, M * D, N)
72+
x = self.dwconv(x)
73+
74+
x = x.reshape(B * M, D, N)
75+
x = self.norm(x)
76+
x = x.reshape(B, M * D, N)
77+
78+
x = self.ffn1_drop1(self.ffn1_pw1(x))
79+
x = self.ffn1_act(x)
80+
x = self.ffn1_drop2(self.ffn1_pw2(x))
81+
x = x.reshape(B, M, D, N)
82+
83+
x = x.permute(0, 2, 1, 3)
84+
x = x.reshape(B, D * M, N)
85+
x = self.ffn2_drop1(self.ffn2_pw1(x))
86+
x = self.ffn2_act(x)
87+
x = self.ffn2_drop2(self.ffn2_pw2(x))
88+
x = x.reshape(B, D, M, N)
89+
x = x.permute(0, 2, 1, 3)
90+
91+
return input_x + x
92+
93+
94+
class ReparamLargeKernelConv(nn.Module):
95+
"""
96+
Reparameterizable Large Kernel Convolution.
97+
98+
This layer uses a large kernel (kernel_size) and
99+
a small kernel in parallel,then adds their outputs.
100+
101+
Parameters
102+
----------
103+
in_channels : int
104+
Number of input channels.
105+
out_channels : int
106+
Number of output channels.
107+
kernel_size : int
108+
Large kernel size.
109+
stride : int
110+
Stride.
111+
groups : int
112+
Number of groups.
113+
small_kernel_size : int
114+
Small kernel size.
115+
"""
116+
117+
def __init__(
118+
self, in_channels, out_channels, kernel_size, stride, groups, small_kernel_size
119+
):
120+
super().__init__()
121+
self.kernel_size = kernel_size
122+
self.small_kernel_size = small_kernel_size
123+
124+
padding = kernel_size // 2
125+
self.lkb_origin = nn.Sequential(
126+
nn.Conv1d(
127+
in_channels,
128+
out_channels,
129+
kernel_size=kernel_size,
130+
stride=stride,
131+
padding=padding,
132+
groups=groups,
133+
bias=False,
134+
),
135+
nn.BatchNorm1d(out_channels),
136+
)
137+
138+
self.small_conv = nn.Sequential(
139+
nn.Conv1d(
140+
in_channels,
141+
out_channels,
142+
kernel_size=small_kernel_size,
143+
stride=stride,
144+
padding=small_kernel_size // 2,
145+
groups=groups,
146+
bias=False,
147+
),
148+
nn.BatchNorm1d(out_channels),
149+
)
150+
151+
def forward(self, x):
152+
return self.lkb_origin(x) + self.small_conv(x)
153+
154+
155+
class Flatten_Head(nn.Module):
156+
"""
157+
Flatten Head.
158+
159+
This layer flattens the input and projects
160+
it to the target window.
161+
162+
Parameters
163+
----------
164+
individual : bool
165+
If True, uses a separate linear projection per variable.
166+
n_vars : int
167+
Number of variables.
168+
nf : int
169+
Number of features.
170+
target_window : int
171+
Length of the target window.
172+
head_dropout : float
173+
Dropout rate.
174+
"""
175+
176+
def __init__(self, individual, n_vars, nf, target_window, head_dropout=0):
177+
super().__init__()
178+
self.individual = individual
179+
self.n_vars = n_vars
180+
181+
if self.individual:
182+
self.linears = nn.ModuleList()
183+
self.dropouts = nn.ModuleList()
184+
self.flattens = nn.ModuleList()
185+
for _ in range(self.n_vars):
186+
self.flattens.append(nn.Flatten(start_dim=-2))
187+
self.linears.append(nn.Linear(nf, target_window))
188+
self.dropouts.append(nn.Dropout(head_dropout))
189+
else:
190+
self.flatten = nn.Flatten(start_dim=-2)
191+
self.linear = nn.Linear(nf, target_window)
192+
self.dropout = nn.Dropout(head_dropout)
23193

24194
def forward(self, x):
25-
residual = x
26-
out = self.act(self.norm(self.dwconv(x)))
27-
out = self.drop(self.act(self.pw_conv1(out)))
28-
out = self.drop(self.pw_conv2(out))
29-
return out + residual
195+
if self.individual:
196+
x_out = []
197+
for i in range(self.n_vars):
198+
z = self.flattens[i](x[:, i, :, :])
199+
z = self.linears[i](z)
200+
z = self.dropouts[i](z)
201+
x_out.append(z)
202+
x = torch.stack(x_out, dim=1)
203+
else:
204+
x = self.flatten(x)
205+
x = self.linear(x)
206+
x = self.dropout(x)
207+
return x

pytorch_forecasting/models/modern_tcn/_modern_tcn_pkg_v2.py

Lines changed: 13 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -73,13 +73,21 @@ def get_test_train_params(cls):
7373
"patch_size": 4,
7474
"use_revin": False,
7575
},
76+
{
77+
"d_model": 16,
78+
"kernel_size": 7,
79+
"small_kernel_size": 3,
80+
"n_blocks": 1,
81+
"d_ff": 32,
82+
"patch_size": 4,
83+
"individual": True,
84+
"use_revin": False,
85+
},
7686
]
7787

78-
default_dm_cfg = {"max_encoder_length": 8, "max_prediction_length": 2}
79-
8088
for param in params:
81-
current_dm_cfg = param.get("datamodule_cfg", {})
82-
default_dm_cfg.update(current_dm_cfg)
83-
param["datamodule_cfg"] = default_dm_cfg
89+
dm_cfg = {"max_encoder_length": 8, "max_prediction_length": 2}
90+
dm_cfg.update(param.get("datamodule_cfg", {}))
91+
param["datamodule_cfg"] = dm_cfg
8492

8593
return params

pytorch_forecasting/models/modern_tcn/_modern_tcn_v2.py

Lines changed: 60 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,10 @@
88
from torch.optim import Optimizer
99

1010
from pytorch_forecasting.layers import RevIN
11-
from pytorch_forecasting.layers._blocks._modern_tcn_block import ModernTCNBlock
11+
from pytorch_forecasting.layers._blocks._modern_tcn_block import (
12+
Flatten_Head,
13+
ModernTCNBlock,
14+
)
1215
from pytorch_forecasting.models.base._base_model_v2 import BaseModel
1316

1417

@@ -27,17 +30,36 @@ class ModernTCN(BaseModel):
2730
d_model : int
2831
Embedding dimension per patch.
2932
kernel_size : int
30-
Kernel size for the depthwise convolution.
33+
Large kernel size for the reparameterizable depthwise convolution.
34+
small_kernel_size : int
35+
Small kernel size for the parallel branch in ReparamLargeKernelConv.
3136
n_blocks : int
32-
Number of ModernTCN blocks.
37+
Number of ModernTCN encoder blocks.
3338
d_ff : int
34-
Hidden dimension in the pointwise conv FFN.
39+
Hidden dimension in the pointwise conv FFNs.
3540
patch_size : int
36-
Number of time steps per patch.
41+
Number of time steps per patch (controls stem stride).
3742
dropout : float
38-
Dropout rate.
43+
Dropout rate applied inside each encoder block.
44+
head_dropout : float
45+
Dropout rate applied inside the Flatten_Head projection.
46+
individual : bool
47+
If True, uses a separate linear projection per variable in the head.
3948
use_revin : bool
40-
Whether to use RevIN normalization.
49+
Whether to apply Reversible Instance Normalization before the encoder.
50+
logging_metrics : list[nn.Module] or None
51+
Additional metrics to log during training.
52+
optimizer : Optimizer or str or None
53+
Optimizer or name of optimizer (default: ``"adam"``).
54+
optimizer_params : dict or None
55+
Additional keyword arguments passed to the optimizer constructor.
56+
lr_scheduler : str or None
57+
Name of a learning rate scheduler recognised by Lightning.
58+
lr_scheduler_params : dict or None
59+
Additional keyword arguments passed to the scheduler constructor.
60+
metadata : dict or None
61+
Dataset metadata injected by the package layer (encoder lengths,
62+
target dim, etc.).
4163
"""
4264

4365
@classmethod
@@ -54,10 +76,13 @@ def __init__(
5476
loss: nn.Module,
5577
d_model: int = 64,
5678
kernel_size: int = 51,
79+
small_kernel_size: int = 5,
5780
n_blocks: int = 2,
5881
d_ff: int = 256,
5982
patch_size: int = 8,
6083
dropout: float = 0.1,
84+
head_dropout: float = 0.1,
85+
individual: bool = False,
6186
use_revin: bool = True,
6287
logging_metrics: list[nn.Module] | None = None,
6388
optimizer: Optimizer | str | None = "adam",
@@ -96,18 +121,31 @@ def __init__(
96121
if self.use_revin:
97122
self.revin = RevIN(num_features=self.n_channels)
98123

99-
self.patch_embed = nn.Linear(patch_size, d_model)
124+
self.stem = nn.Sequential(
125+
nn.Conv1d(1, d_model, kernel_size=patch_size, stride=patch_size),
126+
nn.BatchNorm1d(d_model),
127+
)
100128

101129
self.blocks = nn.ModuleList(
102130
[
103-
ModernTCNBlock(d_model, kernel_size, d_ff, dropout)
131+
ModernTCNBlock(
132+
d_model,
133+
kernel_size,
134+
small_kernel_size,
135+
d_ff,
136+
self.n_channels,
137+
dropout,
138+
)
104139
for _ in range(n_blocks)
105140
]
106141
)
107142

108-
self.head = nn.Linear(
109-
self.n_patches * d_model,
110-
self.prediction_length * self.n_quantiles,
143+
self.head = Flatten_Head(
144+
individual=individual,
145+
n_vars=self.n_channels,
146+
nf=d_model * self.n_patches,
147+
target_window=self.prediction_length * self.n_quantiles,
148+
head_dropout=head_dropout,
111149
)
112150

113151
def forward(self, x: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]:
@@ -125,21 +163,24 @@ def forward(self, x: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]:
125163
B, L, C = input_data.shape
126164

127165
x_enc = input_data.permute(0, 2, 1)
128-
x_enc = x_enc.reshape(B, C, self.n_patches, self.patch_size)
129-
x_enc = self.patch_embed(x_enc)
130-
x_enc = x_enc.permute(0, 1, 3, 2)
131-
x_enc = x_enc.reshape(B * C, self.d_model, self.n_patches)
166+
x_enc = x_enc.reshape(B * C, 1, L)
167+
x_enc = self.stem(x_enc)
132168

169+
x_enc = x_enc.reshape(B, C, self.d_model, self.n_patches)
133170
for block in self.blocks:
134171
x_enc = block(x_enc)
135172

136-
x_enc = x_enc.reshape(B, C, self.d_model, self.n_patches)
137-
x_enc = x_enc.reshape(B, C, -1)
138-
139173
out = self.head(x_enc)
140174
out = out.reshape(B, C, self.prediction_length, self.n_quantiles)
141175
out = out.permute(0, 2, 1, 3)
142176

177+
if self.use_revin:
178+
out = out.permute(0, 1, 3, 2)
179+
out = out.reshape(B, self.prediction_length * self.n_quantiles, C)
180+
out = self.revin(out, mode="denorm")
181+
out = out.reshape(B, self.prediction_length, self.n_quantiles, C)
182+
out = out.permute(0, 1, 3, 2)
183+
143184
target_indices = list(range(self.n_cont_features, C))
144185
out = out[:, :, target_indices, :]
145186

0 commit comments

Comments
 (0)