MCPcopy Create free account
hub / github.com/BIT-MCS/DRL-eFresh / __init__

Method __init__

utils/spatial_att_github.py:7–107  ·  view source on GitHub ↗

Initialize stateful ConvLSTM cell. Parameters ---------- input_channels : ``int`` Number of channels of input tensor. hidden_channels : ``int`` Number of channels of hidden state. kernel_size : ``int`` Size of the c

(self, input_channels, hidden_channels, kernel_size)

Source from the content-addressed store, hash-verified

5
6class ConvLSTMCell(nn.Module):
7 def __init__(self, input_channels, hidden_channels, kernel_size):
8 """Initialize stateful ConvLSTM cell.
9
10 Parameters
11 ----------
12 input_channels : ``int``
13 Number of channels of input tensor.
14 hidden_channels : ``int``
15 Number of channels of hidden state.
16 kernel_size : ``int``
17 Size of the convolutional kernel.
18
19 Paper
20 -----
21 https://papers.nips.cc/paper/5955-convolutional-lstm-network-a-machine-learning-approach-for-precipitation-nowcasting.pdf
22
23 Referenced code
24 ---------------
25 https://github.com/automan000/Convolution_LSTM_PyTorch/blob/master/convolution_lstm.py
26 """
27 super(ConvLSTMCell, self).__init__()
28
29 assert hidden_channels % 2 == 0
30
31 self.input_channels = input_channels
32 self.hidden_channels = hidden_channels
33 self.kernel_size = kernel_size
34 self.num_features = 4
35
36 self.padding = int((kernel_size - 1) / 2)
37
38 self.Wxi = nn.Conv2d(
39 self.input_channels,
40 self.hidden_channels,
41 self.kernel_size,
42 1,
43 self.padding,
44 bias=True,
45 )
46 self.Whi = nn.Conv2d(
47 self.hidden_channels,
48 self.hidden_channels,
49 self.kernel_size,
50 1,
51 self.padding,
52 bias=False,
53 )
54 self.Wxf = nn.Conv2d(
55 self.input_channels,
56 self.hidden_channels,
57 self.kernel_size,
58 1,
59 self.padding,
60 bias=True,
61 )
62 self.Whf = nn.Conv2d(
63 self.hidden_channels,
64 self.hidden_channels,

Callers

nothing calls this directly

Calls 1

__init__Method · 0.45

Tested by

no test coverage detected