Skip to content

Commit ec7d8be

Browse files
Separate the reusable layers
1 parent f294c00 commit ec7d8be

6 files changed

Lines changed: 141 additions & 121 deletions

File tree

Lines changed: 5 additions & 117 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,14 @@
11
"""
2-
ModernTCN building blocks: ReparamLargeKernelConv, ModernTCNBlock, and Flatten_Head.
2+
ModernTCN Block: For Modern Temporal Convolutional Network
33
"""
44

55
import torch
66
import torch.nn as nn
77

8+
from pytorch_forecasting.layers._convolution._reparam_large_kernel_conv import (
9+
ReparamLargeKernelConv,
10+
)
11+
812

913
class ModernTCNBlock(nn.Module):
1014
"""
@@ -89,119 +93,3 @@ def forward(self, x):
8993
x = x.permute(0, 2, 1, 3)
9094

9195
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)
193-
194-
def forward(self, x):
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
Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,3 @@
1+
from pytorch_forecasting.layers._convolution._reparam_large_kernel_conv import ReparamLargeKernelConv
2+
3+
__all__ = ["ReparamLargeKernelConv"]
Lines changed: 67 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,67 @@
1+
"""
2+
Reparameterizable Large Kernel Convolution.
3+
"""
4+
5+
import torch
6+
import torch.nn as nn
7+
8+
9+
class ReparamLargeKernelConv(nn.Module):
10+
"""
11+
Reparameterizable Large Kernel Convolution.
12+
13+
This layer uses a large kernel (kernel_size) and
14+
a small kernel in parallel,then adds their outputs.
15+
16+
Parameters
17+
----------
18+
in_channels : int
19+
Number of input channels.
20+
out_channels : int
21+
Number of output channels.
22+
kernel_size : int
23+
Large kernel size.
24+
stride : int
25+
Stride.
26+
groups : int
27+
Number of groups.
28+
small_kernel_size : int
29+
Small kernel size.
30+
"""
31+
32+
def __init__(
33+
self, in_channels, out_channels, kernel_size, stride, groups, small_kernel_size
34+
):
35+
super().__init__()
36+
self.kernel_size = kernel_size
37+
self.small_kernel_size = small_kernel_size
38+
39+
padding = kernel_size // 2
40+
self.lkb_origin = nn.Sequential(
41+
nn.Conv1d(
42+
in_channels,
43+
out_channels,
44+
kernel_size=kernel_size,
45+
stride=stride,
46+
padding=padding,
47+
groups=groups,
48+
bias=False,
49+
),
50+
nn.BatchNorm1d(out_channels),
51+
)
52+
53+
self.small_conv = nn.Sequential(
54+
nn.Conv1d(
55+
in_channels,
56+
out_channels,
57+
kernel_size=small_kernel_size,
58+
stride=stride,
59+
padding=small_kernel_size // 2,
60+
groups=groups,
61+
bias=False,
62+
),
63+
nn.BatchNorm1d(out_channels),
64+
)
65+
66+
def forward(self, x):
67+
return self.lkb_origin(x) + self.small_conv(x)
Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,3 @@
1+
from pytorch_forecasting.layers._head._flatten_head import Flatten_Head
2+
3+
__all__ = ["Flatten_Head"]
Lines changed: 61 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,61 @@
1+
"""
2+
Flatten Head layer for time series forecasting models.
3+
"""
4+
5+
import torch
6+
import torch.nn as nn
7+
8+
9+
class Flatten_Head(nn.Module):
10+
"""
11+
Flatten Head.
12+
13+
This layer flattens the input and projects
14+
it to the target window.
15+
16+
Parameters
17+
----------
18+
individual : bool
19+
If True, uses a separate linear projection per variable.
20+
n_vars : int
21+
Number of variables.
22+
nf : int
23+
Number of features.
24+
target_window : int
25+
Length of the target window.
26+
head_dropout : float
27+
Dropout rate.
28+
"""
29+
30+
def __init__(self, individual, n_vars, nf, target_window, head_dropout=0):
31+
super().__init__()
32+
self.individual = individual
33+
self.n_vars = n_vars
34+
35+
if self.individual:
36+
self.linears = nn.ModuleList()
37+
self.dropouts = nn.ModuleList()
38+
self.flattens = nn.ModuleList()
39+
for _ in range(self.n_vars):
40+
self.flattens.append(nn.Flatten(start_dim=-2))
41+
self.linears.append(nn.Linear(nf, target_window))
42+
self.dropouts.append(nn.Dropout(head_dropout))
43+
else:
44+
self.flatten = nn.Flatten(start_dim=-2)
45+
self.linear = nn.Linear(nf, target_window)
46+
self.dropout = nn.Dropout(head_dropout)
47+
48+
def forward(self, x):
49+
if self.individual:
50+
x_out = []
51+
for i in range(self.n_vars):
52+
z = self.flattens[i](x[:, i, :, :])
53+
z = self.linears[i](z)
54+
z = self.dropouts[i](z)
55+
x_out.append(z)
56+
x = torch.stack(x_out, dim=1)
57+
else:
58+
x = self.flatten(x)
59+
x = self.linear(x)
60+
x = self.dropout(x)
61+
return x

pytorch_forecasting/models/modern_tcn/_modern_tcn_v2.py

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

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

1715

0 commit comments

Comments
 (0)