MCPcopy Create free account
hub / github.com/Anoise/WTFlib / MultiWaveletCross

Class MultiWaveletCross

LDPS_Graph/layers/MultiWaveletCorrelation.py:61–209  ·  view source on GitHub ↗

1D Multiwavelet Cross Attention layer.

Source from the content-addressed store, hash-verified

59
60
61class MultiWaveletCross(nn.Module):
62 """
63 1D Multiwavelet Cross Attention layer.
64 """
65 def __init__(self, in_channels, out_channels, seq_len_q, seq_len_kv, modes, c=64,
66 k=8, ich=512,
67 L=0,
68 base='legendre',
69 mode_select_method='random',
70 initializer=None, activation='tanh',
71 **kwargs):
72 super(MultiWaveletCross, self).__init__()
73 print('base', base)
74
75 self.c = c
76 self.k = k
77 self.L = L
78 H0, H1, G0, G1, PHI0, PHI1 = get_filter(base, k)
79 H0r = H0 @ PHI0
80 G0r = G0 @ PHI0
81 H1r = H1 @ PHI1
82 G1r = G1 @ PHI1
83
84 H0r[np.abs(H0r) < 1e-8] = 0
85 H1r[np.abs(H1r) < 1e-8] = 0
86 G0r[np.abs(G0r) < 1e-8] = 0
87 G1r[np.abs(G1r) < 1e-8] = 0
88 self.max_item = 3
89
90 self.attn1 = FourierCrossAttentionW(in_channels=in_channels, out_channels=out_channels, seq_len_q=seq_len_q,
91 seq_len_kv=seq_len_kv, modes=modes, activation=activation,
92 mode_select_method=mode_select_method)
93 self.attn2 = FourierCrossAttentionW(in_channels=in_channels, out_channels=out_channels, seq_len_q=seq_len_q,
94 seq_len_kv=seq_len_kv, modes=modes, activation=activation,
95 mode_select_method=mode_select_method)
96 self.attn3 = FourierCrossAttentionW(in_channels=in_channels, out_channels=out_channels, seq_len_q=seq_len_q,
97 seq_len_kv=seq_len_kv, modes=modes, activation=activation,
98 mode_select_method=mode_select_method)
99 self.attn4 = FourierCrossAttentionW(in_channels=in_channels, out_channels=out_channels, seq_len_q=seq_len_q,
100 seq_len_kv=seq_len_kv, modes=modes, activation=activation,
101 mode_select_method=mode_select_method)
102 self.T0 = nn.Linear(k, k)
103 self.register_buffer('ec_s', torch.Tensor(
104 np.concatenate((H0.T, H1.T), axis=0)))
105 self.register_buffer('ec_d', torch.Tensor(
106 np.concatenate((G0.T, G1.T), axis=0)))
107
108 self.register_buffer('rc_e', torch.Tensor(
109 np.concatenate((H0r, G0r), axis=0)))
110 self.register_buffer('rc_o', torch.Tensor(
111 np.concatenate((H1r, G1r), axis=0)))
112
113 self.Lk = nn.Linear(ich, c * k)
114 self.Lq = nn.Linear(ich, c * k)
115 self.Lv = nn.Linear(ich, c * k)
116 self.out = nn.Linear(c * k, ich)
117 self.modes1 = modes
118

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected