MCPcopy Create free account
hub / github.com/BorealisAI/scaleformer / __init__

Method __init__

models/FiLM.py:125–164  ·  view source on GitHub ↗
(self, in_channels, out_channels,seq_len, modes1,compression=0,ratio=0.5,mode_type=0)

Source from the content-addressed store, hash-verified

123
124class SpectralConv1d(nn.Module):
125 def __init__(self, in_channels, out_channels,seq_len, modes1,compression=0,ratio=0.5,mode_type=0):
126 super(SpectralConv1d, self).__init__()
127
128
129 """
130 1D Fourier layer. It does FFT, linear transform, and Inverse FFT.
131 """
132
133 self.in_channels = in_channels
134 self.out_channels = out_channels
135 self.modes1 = modes1 #Number of Fourier modes to multiply, at most floor(N/2) + 1
136 self.compression = compression
137 self.ratio = ratio
138 self.mode_type=mode_type
139 if self.mode_type ==1:
140 modes2 = modes1
141 self.modes2 =min(modes2,seq_len//2)
142 self.index0 = list(range(0, int(ratio*min(seq_len//2, modes2))))
143 self.index1 = list(range(len(self.index0),self.modes2))
144 np.random.shuffle(self.index1)
145 self.index1 = self.index1[:min(seq_len//2,self.modes2)-int(ratio*min(seq_len//2, modes2))]
146 self.index = self.index0+self.index1
147 self.index.sort()
148 elif self.mode_type > 1:
149 modes2 = modes1
150 self.modes2 =min(modes2,seq_len//2)
151 self.index = list(range(0, seq_len//2))
152 np.random.shuffle(self.index)
153 self.index = self.index[:self.modes2]
154 else:
155 self.modes2 =min(modes1,seq_len//2)
156 self.index = list(range(0, self.modes2))
157
158 self.scale = (1 / (in_channels*out_channels))
159 self.weights1 = nn.Parameter(self.scale * torch.rand(in_channels, out_channels, len(self.index), dtype=torch.cfloat))
160 if self.compression > 0:
161 print('compressed version')
162 self.weights0 = nn.Parameter(self.scale * torch.rand(in_channels,self.compression,dtype=torch.cfloat))
163 self.weights1 = nn.Parameter(self.scale * torch.rand(self.compression,self.compression, len(self.index), dtype=torch.cfloat))
164 self.weights2 = nn.Parameter(self.scale * torch.rand(self.compression,out_channels, dtype=torch.cfloat))
165
166
167 def forward(self, x):

Callers

nothing calls this directly

Calls 1

__init__Method · 0.45

Tested by

no test coverage detected