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.
| 9 | |
| 10 | |
| 11 | class 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 |
nothing calls this directly
no outgoing calls
no test coverage detected