MCPcopy Create free account
hub / github.com/kwuking/TimeMixer / AutoCorrelation

Class AutoCorrelation

layers/AutoCorrelation.py:11–128  ·  view source on GitHub ↗

AutoCorrelation Mechanism with the following two phases: (1) period-based dependencies discovery (2) time delay aggregation This block can replace the self-attention family mechanism seamlessly.

Source from the content-addressed store, hash-verified

9
10
11class AutoCorrelation(nn.Module):
12 """
13 AutoCorrelation Mechanism with the following two phases:
14 (1) period-based dependencies discovery
15 (2) time delay aggregation
16 This block can replace the self-attention family mechanism seamlessly.
17 """
18
19 def __init__(self, mask_flag=True, factor=1, scale=None, attention_dropout=0.1, output_attention=False):
20 super(AutoCorrelation, self).__init__()
21 self.factor = factor
22 self.scale = scale
23 self.mask_flag = mask_flag
24 self.output_attention = output_attention
25 self.dropout = nn.Dropout(attention_dropout)
26
27 def time_delay_agg_training(self, values, corr):
28 """
29 SpeedUp version of Autocorrelation (a batch-normalization style design)
30 This is for the training phase.
31 """
32 head = values.shape[1]
33 channel = values.shape[2]
34 length = values.shape[3]
35 # find top k
36 top_k = int(self.factor * math.log(length))
37 mean_value = torch.mean(torch.mean(corr, dim=1), dim=1)
38 index = torch.topk(torch.mean(mean_value, dim=0), top_k, dim=-1)[1]
39 weights = torch.stack([mean_value[:, index[i]] for i in range(top_k)], dim=-1)
40 # update corr
41 tmp_corr = torch.softmax(weights, dim=-1)
42 # aggregation
43 tmp_values = values
44 delays_agg = torch.zeros_like(values).float()
45 for i in range(top_k):
46 pattern = torch.roll(tmp_values, -int(index[i]), -1)
47 delays_agg = delays_agg + pattern * \
48 (tmp_corr[:, i].unsqueeze(1).unsqueeze(1).unsqueeze(1).repeat(1, head, channel, length))
49 return delays_agg
50
51 def time_delay_agg_inference(self, values, corr):
52 """
53 SpeedUp version of Autocorrelation (a batch-normalization style design)
54 This is for the inference phase.
55 """
56 batch = values.shape[0]
57 head = values.shape[1]
58 channel = values.shape[2]
59 length = values.shape[3]
60 # index init
61 init_index = torch.arange(length).unsqueeze(0).unsqueeze(0).unsqueeze(0).repeat(batch, head, channel, 1).cuda()
62 # find top k
63 top_k = int(self.factor * math.log(length))
64 mean_value = torch.mean(torch.mean(corr, dim=1), dim=1)
65 weights, delay = torch.topk(mean_value, top_k, dim=-1)
66 # update corr
67 tmp_corr = torch.softmax(weights, dim=-1)
68 # aggregation

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected