|
1 | 1 | """ |
2 | | -ModernTCN building blocks: ReparamLargeKernelConv, ModernTCNBlock, and Flatten_Head. |
| 2 | +ModernTCN Block: For Modern Temporal Convolutional Network |
3 | 3 | """ |
4 | 4 |
|
5 | 5 | import torch |
6 | 6 | import torch.nn as nn |
7 | 7 |
|
| 8 | +from pytorch_forecasting.layers._convolution._reparam_large_kernel_conv import ( |
| 9 | + ReparamLargeKernelConv, |
| 10 | +) |
| 11 | + |
8 | 12 |
|
9 | 13 | class ModernTCNBlock(nn.Module): |
10 | 14 | """ |
@@ -89,119 +93,3 @@ def forward(self, x): |
89 | 93 | x = x.permute(0, 2, 1, 3) |
90 | 94 |
|
91 | 95 | 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 |
0 commit comments