Coverage for hopwise/model/sequential_recommender/fearec.py: 71%
442 statements
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 13:25 +0000
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 13:25 +0000
1# @Time : 2023/10/27
2# @Author : Kesha Ou
3# @Email : keishaou@gmail.com
5r"""FEARec
6################################################
8Reference:
9 Xinyu Du et al. "Frequency Enhanced Hybrid Attention Network for Sequential Recommendation."
10 In SIGIR 2023.
12Reference code:
13 https://github.com/sudaada/FEARec
15"""
17import math
18import random
20import numpy as np
21import torch
22import torch.nn.functional as F
23import torch.nn.functional as fn
24from torch import nn
26from hopwise.data.interaction import Interaction
27from hopwise.model.abstract_recommender import SequentialRecommender
28from hopwise.model.loss import BPRLoss
31class FEARec(SequentialRecommender):
32 def __init__(self, config, dataset):
33 super().__init__(config, dataset)
35 # load parameters info
36 self.dataset = dataset
37 self.config = config
38 self.n_layers = config["n_layers"]
39 self.n_heads = config["n_heads"]
40 self.hidden_size = config["hidden_size"] # same as embedding_size
41 self.inner_size = config["inner_size"] # the dimensionality in feed-forward layer
42 self.hidden_dropout_prob = config["hidden_dropout_prob"]
43 self.attn_dropout_prob = config["attn_dropout_prob"]
44 self.hidden_act = config["hidden_act"]
45 self.layer_norm_eps = config["layer_norm_eps"]
47 self.lmd = config["lmd"]
48 self.lmd_sem = config["lmd_sem"]
50 self.initializer_range = config["initializer_range"]
51 self.loss_type = config["loss_type"]
52 self.same_item_index = self.get_same_item_index(dataset)
54 # define layers and loss
55 self.item_embedding = nn.Embedding(self.n_items, self.hidden_size, padding_idx=0)
56 self.position_embedding = nn.Embedding(self.max_seq_length, self.hidden_size)
57 self.item_encoder = FEAEncoder(
58 n_layers=self.n_layers,
59 n_heads=self.n_heads,
60 hidden_size=self.hidden_size,
61 inner_size=self.inner_size,
62 hidden_dropout_prob=self.hidden_dropout_prob,
63 attn_dropout_prob=self.attn_dropout_prob,
64 hidden_act=self.hidden_act,
65 layer_norm_eps=self.layer_norm_eps,
66 config=self.config,
67 )
69 self.LayerNorm = nn.LayerNorm(self.hidden_size, eps=self.layer_norm_eps)
70 self.dropout = nn.Dropout(self.hidden_dropout_prob)
72 if self.loss_type == "BPR":
73 self.loss_fct = BPRLoss()
74 elif self.loss_type == "CE":
75 self.loss_fct = nn.CrossEntropyLoss()
76 else:
77 raise NotImplementedError("Make sure 'loss_type' in ['BPR', 'CE']!")
79 self.ssl = config["contrast"]
80 self.tau = config["tau"]
81 self.sim = config["sim"]
82 self.fredom = config["fredom"]
83 self.fredom_type = config["fredom_type"]
84 self.batch_size = config["train_batch_size"]
85 self.mask_default = self.mask_correlated_samples(batch_size=self.batch_size)
86 self.aug_nce_fct = nn.CrossEntropyLoss()
87 self.sem_aug_nce_fct = nn.CrossEntropyLoss()
89 # parameters initialization
90 self.apply(self._init_weights)
92 def get_same_item_index(self, dataset):
93 same_target_index = {}
94 target_item = dataset.inter_feat[self.ITEM_ID].numpy()
96 for index, item_id in enumerate(target_item):
97 all_index_same_id = np.where(target_item == item_id)[0]
98 same_target_index[item_id] = all_index_same_id
100 return same_target_index
102 def _init_weights(self, module):
103 """Initialize the weights"""
104 if isinstance(module, (nn.Linear, nn.Embedding)):
105 # Slightly different from the TF version which uses truncated_normal for initialization
106 # cf https://github.com/pytorch/pytorch/pull/5617
107 module.weight.data.normal_(mean=0.0, std=self.initializer_range)
108 # module.weight.data = self.truncated_normal_(tensor=module.weight.data, mean=0, std=self.initializer_range) # noqa: E501
109 elif isinstance(module, nn.LayerNorm):
110 module.bias.data.zero_()
111 module.weight.data.fill_(1.0)
112 if isinstance(module, nn.Linear) and module.bias is not None:
113 module.bias.data.zero_()
115 def truncated_normal_(self, tensor, mean=0, std=0.09):
116 with torch.no_grad():
117 size = tensor.shape
118 tmp = tensor.new_empty(size + (4,)).normal_()
119 valid = (tmp < 2) & (tmp > -2) # noqa: PLR2004
120 ind = valid.max(-1, keepdim=True)[1]
121 tensor.data.copy_(tmp.gather(-1, ind).squeeze(-1))
122 tensor.data.mul_(std).add_(mean)
123 return tensor
125 def get_attention_mask(self, item_seq):
126 """Generate left-to-right uni-directional attention mask for multi-head attention."""
127 attention_mask = (item_seq > 0).long()
128 extended_attention_mask = attention_mask.unsqueeze(1).unsqueeze(2) # torch.int64
129 # mask for left-to-right unidirectional
130 max_len = attention_mask.size(-1)
131 attn_shape = (1, max_len, max_len)
132 subsequent_mask = torch.triu(torch.ones(attn_shape), diagonal=1) # torch.uint8
133 subsequent_mask = (subsequent_mask == 0).unsqueeze(1)
134 subsequent_mask = subsequent_mask.long().to(item_seq.device)
136 extended_attention_mask = extended_attention_mask * subsequent_mask
137 extended_attention_mask = extended_attention_mask.to(dtype=next(self.parameters()).dtype) # fp16 compatibility
138 extended_attention_mask = (1.0 - extended_attention_mask) * -10000.0
139 return extended_attention_mask
141 def get_bi_attention_mask(self, item_seq):
142 """Generate bidirectional attention mask for multi-head attention."""
143 attention_mask = (item_seq > 0).long()
144 extended_attention_mask = attention_mask.unsqueeze(1).unsqueeze(2) # torch.int64
145 # bidirectional mask
146 extended_attention_mask = extended_attention_mask.to(dtype=next(self.parameters()).dtype) # fp16 compatibility
147 extended_attention_mask = (1.0 - extended_attention_mask) * -10000.0
148 return extended_attention_mask
150 def forward(self, item_seq, item_seq_len):
151 position_ids = torch.arange(item_seq.size(1), dtype=torch.long, device=item_seq.device)
152 position_ids = position_ids.unsqueeze(0).expand_as(item_seq)
153 position_embedding = self.position_embedding(position_ids)
155 item_emb = self.item_embedding(item_seq)
156 input_emb = item_emb + position_embedding
157 input_emb = self.LayerNorm(input_emb)
158 input_emb = self.dropout(input_emb)
160 extended_attention_mask = self.get_attention_mask(item_seq)
161 # extended_attention_mask = self.get_bi_attention_mask(item_seq)
163 trm_output = self.item_encoder(input_emb, extended_attention_mask, output_all_encoded_layers=True)
164 output = trm_output[-1]
165 output = self.gather_indexes(output, item_seq_len - 1)
167 return output # [B H]
169 @staticmethod
170 def alignment(x, y):
171 x, y = F.normalize(x, dim=-1), F.normalize(y, dim=-1)
172 return (x - y).norm(p=2, dim=1).pow(2).mean()
174 @staticmethod
175 def uniformity(x):
176 x = F.normalize(x, dim=-1)
177 x = abs(x)
178 return torch.pdist(x, p=2).pow(2).mul(-2).exp().mean().log()
180 def calculate_loss(self, interaction):
181 same_item_index = self.same_item_index
182 sem_pos_lengths = []
183 sem_pos_seqs = []
184 dataset = self.dataset
185 target_items = interaction[self.ITEM_ID]
186 for i, target_item_id in enumerate(target_items):
187 item_id = target_item_id.item()
188 targets_index = same_item_index[item_id]
189 lens = len(targets_index)
190 if lens == 0:
191 print("error")
192 remaining_indices = targets_index.copy()
193 while len(remaining_indices) > 0:
194 sample_index = random.choice(remaining_indices)
195 remaining_indices = remaining_indices[remaining_indices != sample_index]
196 cur_item_list = interaction[self.ITEM_SEQ][i].to("cpu")
197 sample_item_list = dataset.inter_feat[self.ITEM_SEQ][sample_index]
198 are_equal = torch.equal(cur_item_list, sample_item_list)
199 sample_item_length = dataset.inter_feat[self.ITEM_SEQ_LEN][sample_index]
201 if not are_equal or len(remaining_indices) == 0:
202 sem_pos_lengths.append(sample_item_length)
203 sem_pos_seqs.append(sample_item_list)
204 break
206 sem_pos_lengths = torch.stack(sem_pos_lengths).to(self.device)
207 sem_pos_seqs = torch.stack(sem_pos_seqs).to(self.device)
209 interaction.update(Interaction({"sem_aug": sem_pos_seqs, "sem_aug_lengths": sem_pos_lengths}))
211 item_seq = interaction[self.ITEM_SEQ]
212 item_seq_len = interaction[self.ITEM_SEQ_LEN]
213 seq_output = self.forward(item_seq, item_seq_len)
214 pos_items = interaction[self.POS_ITEM_ID]
215 if self.loss_type == "BPR":
216 neg_items = interaction[self.NEG_ITEM_ID]
217 pos_items_emb = self.item_embedding(pos_items)
218 neg_items_emb = self.item_embedding(neg_items)
219 pos_score = torch.sum(seq_output * pos_items_emb, dim=-1) # [B]
220 neg_score = torch.sum(seq_output * neg_items_emb, dim=-1) # [B]
221 loss = self.loss_fct(pos_score, neg_score)
222 else: # self.loss_type = 'CE'
223 test_item_emb = self.item_embedding.weight
224 logits = torch.matmul(seq_output, test_item_emb.transpose(0, 1))
225 loss = self.loss_fct(logits, pos_items)
227 # Unsupervised NCE
228 if self.ssl in ["us", "un"]:
229 aug_seq_output = self.forward(item_seq, item_seq_len)
230 nce_logits, nce_labels = self.info_nce(
231 seq_output,
232 aug_seq_output,
233 temp=self.tau,
234 batch_size=item_seq_len.shape[0],
235 sim=self.sim,
236 )
238 loss += self.lmd * self.aug_nce_fct(nce_logits, nce_labels)
240 # Supervised NCE
241 if self.ssl in ["us", "su"]:
242 sem_aug, sem_aug_lengths = (
243 interaction["sem_aug"],
244 interaction["sem_aug_lengths"],
245 )
246 sem_aug_seq_output = self.forward(sem_aug, sem_aug_lengths)
248 sem_nce_logits, sem_nce_labels = self.info_nce(
249 seq_output,
250 sem_aug_seq_output,
251 temp=self.tau,
252 batch_size=item_seq_len.shape[0],
253 sim=self.sim,
254 )
256 loss += self.lmd_sem * self.aug_nce_fct(sem_nce_logits, sem_nce_labels)
258 if self.ssl == "us_x":
259 aug_seq_output = self.forward(item_seq, item_seq_len)
260 sem_aug, sem_aug_lengths = (
261 interaction["sem_aug"],
262 interaction["sem_aug_lengths"],
263 )
265 sem_aug_seq_output = self.forward(sem_aug, sem_aug_lengths)
266 sem_nce_logits, sem_nce_labels = self.info_nce(
267 aug_seq_output,
268 sem_aug_seq_output,
269 temp=self.tau,
270 batch_size=item_seq_len.shape[0],
271 sim=self.sim,
272 )
274 loss += self.lmd_sem * self.aug_nce_fct(sem_nce_logits, sem_nce_labels)
276 # frequency domain loss
277 if self.fredom:
278 seq_output_f = torch.fft.rfft(seq_output, dim=1, norm="ortho")
279 aug_seq_output_f = torch.fft.rfft(aug_seq_output, dim=1, norm="ortho")
280 sem_aug_seq_output_f = torch.fft.rfft(sem_aug_seq_output, dim=1, norm="ortho")
281 if self.fredom_type in ["us", "un"]:
282 loss += 0.1 * abs(seq_output_f - aug_seq_output_f).flatten().mean()
283 if self.fredom_type in ["us", "su"]:
284 loss += 0.1 * abs(seq_output_f - sem_aug_seq_output_f).flatten().mean()
285 if self.fredom_type == "us_x":
286 loss += 0.1 * abs(aug_seq_output_f - sem_aug_seq_output_f).flatten().mean()
288 return loss
290 def mask_correlated_samples(self, batch_size):
291 N = 2 * batch_size
292 mask = torch.ones((N, N), dtype=bool)
293 mask = mask.fill_diagonal_(0)
294 for i in range(batch_size):
295 mask[i, batch_size + i] = 0
296 mask[batch_size + i, i] = 0
297 return mask
299 def info_nce(self, z_i, z_j, temp, batch_size, sim="dot"):
300 """We do not sample negative examples explicitly.
301 Instead, given a positive pair, similar to (Chen et al., 2017), we treat the other 2(N − 1) augmented examples within a minibatch as negative examples.
302 """ # noqa: E501
303 N = 2 * batch_size
305 z = torch.cat((z_i, z_j), dim=0)
307 if sim == "cos":
308 sim = nn.functional.cosine_similarity(z.unsqueeze(1), z.unsqueeze(0), dim=2) / temp
309 elif sim == "dot":
310 sim = torch.mm(z, z.T) / temp
312 sim_i_j = torch.diag(sim, batch_size)
313 sim_j_i = torch.diag(sim, -batch_size)
315 positive_samples = torch.cat((sim_i_j, sim_j_i), dim=0).reshape(N, 1)
316 if batch_size != self.batch_size:
317 mask = self.mask_correlated_samples(batch_size)
318 else:
319 mask = self.mask_default
320 negative_samples = sim[mask].reshape(N, -1)
322 labels = torch.zeros(N).to(positive_samples.device).long()
323 logits = torch.cat((positive_samples, negative_samples), dim=1)
324 return logits, labels
326 def decompose(self, z_i, z_j, origin_z, batch_size):
327 """We do not sample negative examples explicitly.
328 Instead, given a positive pair, similar to (Chen et al., 2017), we treat the other 2(N − 1) augmented examples within a minibatch as negative examples.
329 """ # noqa: E501
330 N = 2 * batch_size
332 z = torch.cat((z_i, z_j), dim=0)
334 # pairwise l2 distace
335 sim = torch.cdist(z, z, p=2)
337 sim_i_j = torch.diag(sim, batch_size)
338 sim_j_i = torch.diag(sim, -batch_size)
340 positive_samples = torch.cat((sim_i_j, sim_j_i), dim=0).reshape(N, 1)
341 alignment = positive_samples.mean()
343 # pairwise l2 distace
344 sim = torch.cdist(origin_z, origin_z, p=2)
345 mask = torch.ones((batch_size, batch_size), dtype=bool)
346 mask = mask.fill_diagonal_(0)
347 negative_samples = sim[mask].reshape(batch_size, -1)
348 uniformity = torch.log(torch.exp(-2 * negative_samples).mean())
350 return alignment, uniformity
352 def predict(self, interaction):
353 item_seq = interaction[self.ITEM_SEQ]
354 item_seq_len = interaction[self.ITEM_SEQ_LEN]
355 test_item = interaction[self.ITEM_ID]
356 seq_output = self.forward(item_seq, item_seq_len)
357 test_item_emb = self.item_embedding(test_item)
358 scores = torch.mul(seq_output, test_item_emb).sum(dim=1) # [B]
359 return scores
361 def full_sort_predict(self, interaction):
362 item_seq = interaction[self.ITEM_SEQ]
363 item_seq_len = interaction[self.ITEM_SEQ_LEN]
364 seq_output = self.forward(item_seq, item_seq_len)
365 test_items_emb = self.item_embedding.weight
366 scores = torch.matmul(seq_output, test_items_emb.transpose(0, 1)) # [B n_items]
367 return scores
370class HybridAttention(nn.Module):
371 """Hybrid Attention layer: combine time domain self-attention layer and frequency domain attention layer.
373 Args:
374 input_tensor (torch.Tensor): the input of the multi-head Hybrid Attention layer
375 attention_mask (torch.Tensor): the attention mask for input tensor
377 Returns:
378 hidden_states (torch.Tensor): the output of the multi-head Hybrid Attention layer
380 """
382 def __init__(
383 self,
384 n_heads,
385 hidden_size,
386 hidden_dropout_prob,
387 attn_dropout_prob,
388 layer_norm_eps,
389 i,
390 config,
391 ):
392 super().__init__()
393 if hidden_size % n_heads != 0:
394 raise ValueError(
395 "The hidden size (%d) is not a multiple of the number of attention heads (%d)" % (hidden_size, n_heads)
396 )
398 self.factor = config["topk_factor"]
399 self.scale = None
400 self.mask_flag = True
401 self.output_attention = False
402 self.dropout = nn.Dropout(0.1)
403 self.config = config
404 self.num_attention_heads = n_heads
405 self.attention_head_size = int(hidden_size / n_heads)
406 self.all_head_size = self.num_attention_heads * self.attention_head_size
407 self.query_layer = nn.Linear(hidden_size, self.all_head_size)
408 self.key_layer = nn.Linear(hidden_size, self.all_head_size)
409 self.value_layer = nn.Linear(hidden_size, self.all_head_size)
410 self.attn_dropout = nn.Dropout(attn_dropout_prob)
411 self.dense = nn.Linear(hidden_size, hidden_size)
412 self.LayerNorm = nn.LayerNorm(hidden_size, eps=layer_norm_eps)
413 self.out_dropout = nn.Dropout(hidden_dropout_prob)
414 self.filter_mixer = None
415 self.global_ratio = config["global_ratio"]
416 self.n_layers = config["n_layers"]
417 if self.global_ratio > (1 / self.n_layers):
418 print(f"{self.global_ratio}>{1 / self.n_layers}:{self.global_ratio > (1 / self.n_layers)}")
419 self.filter_mixer = "G"
420 else:
421 print(f"{self.global_ratio}>{1 / self.n_layers}:{self.global_ratio > (1 / self.n_layers)}")
422 self.filter_mixer = "L"
423 self.max_item_list_length = config["MAX_ITEM_LIST_LENGTH"]
424 self.dual_domain = config["dual_domain"]
425 self.slide_step = ((self.max_item_list_length // 2 + 1) * (1 - self.global_ratio)) // (self.n_layers - 1)
426 self.local_ratio = 1 / self.n_layers
427 self.filter_size = self.local_ratio * (self.max_item_list_length // 2 + 1)
429 if self.filter_mixer == "G":
430 self.w = self.global_ratio
431 self.s = self.slide_step
433 if self.filter_mixer == "L":
434 self.w = self.local_ratio
435 self.s = self.filter_size
437 self.left = int(((self.max_item_list_length // 2 + 1) * (1 - self.w)) - (i * self.s))
438 self.right = int((self.max_item_list_length // 2 + 1) - i * self.s)
440 self.q_index = list(range(self.left, self.right))
441 self.k_index = list(range(self.left, self.right))
442 self.v_index = list(range(self.left, self.right))
443 # if sample in time domain
444 self.std = config["std"]
445 if self.std:
446 self.time_q_index = self.q_index
447 self.time_k_index = self.k_index
448 self.time_v_index = self.v_index
449 else:
450 self.time_q_index = list(range(self.max_item_list_length // 2 + 1))
451 self.time_k_index = list(range(self.max_item_list_length // 2 + 1))
452 self.time_v_index = list(range(self.max_item_list_length // 2 + 1))
454 print(f"modes_q={len(self.q_index)}, index_q={self.q_index}")
455 print(f"modes_k={len(self.k_index)}, index_k={self.k_index}")
456 print(f"modes_v={len(self.v_index)}, index_v={self.v_index}")
458 if self.config["dual_domain"]:
459 self.spatial_ratio = self.config["spatial_ratio"]
461 def transpose_for_scores(self, x):
462 new_x_shape = x.size()[:-1] + (
463 self.num_attention_heads,
464 self.attention_head_size,
465 )
466 x = x.view(*new_x_shape)
467 # return x.permute(0, 2, 1, 3)
468 return x
470 def time_delay_agg_training(self, values, corr):
471 """SpeedUp version of Autocorrelation (a batch-normalization style design)
472 This is for the training phase.
473 """
474 head = values.shape[1]
475 channel = values.shape[2]
476 length = values.shape[3]
477 # find top k
478 top_k = int(self.factor * math.log(length))
479 mean_value = torch.mean(torch.mean(corr, dim=1), dim=1)
480 index = torch.topk(torch.mean(mean_value, dim=0), top_k, dim=-1)[1]
481 weights = torch.stack([mean_value[:, index[i]] for i in range(top_k)], dim=-1)
482 # update corr
483 tmp_corr = torch.softmax(weights, dim=-1)
484 # aggregation
485 tmp_values = values
486 delays_agg = torch.zeros_like(values).float()
487 for i in range(top_k):
488 pattern = torch.roll(tmp_values, -int(index[i]), -1)
489 delays_agg = delays_agg + pattern * (
490 tmp_corr[:, i].unsqueeze(1).unsqueeze(1).unsqueeze(1).repeat(1, head, channel, length)
491 )
492 return delays_agg
494 def time_delay_agg_inference(self, values, corr):
495 """SpeedUp version of Autocorrelation (a batch-normalization style design)
496 This is for the inference phase.
497 """
498 batch = values.shape[0]
499 head = values.shape[1]
500 channel = values.shape[2]
501 length = values.shape[3]
502 # index init
503 init_index = (
504 torch.arange(length)
505 .unsqueeze(0)
506 .unsqueeze(0)
507 .unsqueeze(0)
508 .repeat(batch, head, channel, 1)
509 .to(values.device)
510 )
511 # find top k
512 top_k = int(self.factor * math.log(length))
513 mean_value = torch.mean(torch.mean(corr, dim=1), dim=1)
514 weights, delay = torch.topk(mean_value, top_k, dim=-1)
515 # update corr
516 tmp_corr = torch.softmax(weights, dim=-1)
517 # aggregation
518 tmp_values = values.repeat(1, 1, 1, 2)
519 delays_agg = torch.zeros_like(values).float()
520 for i in range(top_k):
521 tmp_delay = init_index + delay[:, i].unsqueeze(1).unsqueeze(1).unsqueeze(1).repeat(
522 1, head, channel, length
523 )
524 pattern = torch.gather(tmp_values, dim=-1, index=tmp_delay)
525 delays_agg = delays_agg + pattern * (
526 tmp_corr[:, i].unsqueeze(1).unsqueeze(1).unsqueeze(1).repeat(1, head, channel, length)
527 )
528 return delays_agg
530 def forward(self, input_tensor, attention_mask):
531 mixed_query_layer = self.query_layer(input_tensor)
532 mixed_key_layer = self.key_layer(input_tensor)
533 mixed_value_layer = self.value_layer(input_tensor)
535 queries = self.transpose_for_scores(mixed_query_layer)
536 keys = self.transpose_for_scores(mixed_key_layer)
537 values = self.transpose_for_scores(mixed_value_layer)
539 # B, H, L, E = query_layer.shape
540 # AutoFormer
541 B, L, H, E = queries.shape
542 _, S, _, D = values.shape
543 if L > S:
544 zeros = torch.zeros_like(queries[:, : (L - S), :]).float()
545 values = torch.cat([values, zeros], dim=1)
546 keys = torch.cat([keys, zeros], dim=1)
547 else:
548 values = values[:, :L, :, :]
549 keys = keys[:, :L, :, :]
551 # period-based dependencies
552 q_fft = torch.fft.rfft(queries.permute(0, 2, 3, 1).contiguous(), dim=-1)
553 k_fft = torch.fft.rfft(keys.permute(0, 2, 3, 1).contiguous(), dim=-1)
555 # put into an empty box for sampling
556 q_fft_box = torch.zeros(B, H, E, len(self.q_index), device=q_fft.device, dtype=torch.cfloat)
557 q_fft_box = q_fft[:, :, :, self.q_index]
559 k_fft_box = torch.zeros(B, H, E, len(self.k_index), device=q_fft.device, dtype=torch.cfloat)
560 k_fft_box = k_fft[:, :, :, self.q_index]
561 res = q_fft_box * torch.conj(k_fft_box)
563 if self.config["use_filter"]:
564 # filter
565 weight = torch.view_as_complex(self.complex_weight)
566 res = res * weight
568 box_res = torch.zeros(B, H, E, L // 2 + 1, device=q_fft.device, dtype=torch.cfloat)
569 box_res[:, :, :, self.q_index] = res
571 corr = torch.fft.irfft(box_res, dim=-1)
573 # time delay agg
574 if self.training:
575 V = self.time_delay_agg_training(values.permute(0, 2, 3, 1).contiguous(), corr).permute(0, 3, 1, 2)
576 else:
577 V = self.time_delay_agg_inference(values.permute(0, 2, 3, 1).contiguous(), corr).permute(0, 3, 1, 2)
579 new_context_layer_shape = V.size()[:-2] + (self.all_head_size,)
580 context_layer = V.view(*new_context_layer_shape)
582 if self.dual_domain:
583 # put into an empty box for sampling
584 # q
585 q_fft_box = torch.zeros(B, H, E, len(self.time_q_index), device=q_fft.device, dtype=torch.cfloat)
586 q_fft_box = q_fft[:, :, :, self.time_q_index]
587 spatial_q = torch.zeros(B, H, E, L // 2 + 1, device=q_fft.device, dtype=torch.cfloat)
588 spatial_q[:, :, :, self.time_q_index] = q_fft_box
590 # k
591 k_fft_box = torch.zeros(B, H, E, len(self.time_k_index), device=q_fft.device, dtype=torch.cfloat)
592 k_fft_box = k_fft[:, :, :, self.time_k_index]
593 spatial_k = torch.zeros(B, H, E, L // 2 + 1, device=k_fft.device, dtype=torch.cfloat)
594 spatial_k[:, :, :, self.time_k_index] = k_fft_box
596 # v
597 v_fft = torch.fft.rfft(values.permute(0, 2, 3, 1).contiguous(), dim=-1)
598 # put into an empty box for sampling
599 v_fft_box = torch.zeros(B, H, E, len(self.time_v_index), device=v_fft.device, dtype=torch.cfloat)
600 v_fft_box = v_fft[:, :, :, self.time_v_index]
601 spatial_v = torch.zeros(B, H, E, L // 2 + 1, device=v_fft.device, dtype=torch.cfloat)
602 spatial_v[:, :, :, self.time_v_index] = v_fft_box
604 queries = torch.fft.irfft(spatial_q, dim=-1)
605 keys = torch.fft.irfft(spatial_k, dim=-1)
606 values = torch.fft.irfft(spatial_v, dim=-1)
608 queries = queries.permute(0, 1, 3, 2)
609 keys = keys.permute(0, 1, 3, 2)
610 values = values.permute(0, 1, 3, 2)
612 attention_scores = torch.matmul(queries, keys.transpose(-1, -2))
613 attention_scores = attention_scores / math.sqrt(self.attention_head_size)
615 attention_scores = attention_scores + attention_mask
616 attention_probs = nn.Softmax(dim=-1)(attention_scores)
617 attention_probs = self.attn_dropout(attention_probs)
618 qkv = torch.matmul(attention_probs, values)
619 context_layer_spatial = qkv.permute(0, 2, 1, 3).contiguous()
620 new_context_layer_shape = context_layer_spatial.size()[:-2] + (self.all_head_size,)
621 context_layer_spatial = context_layer_spatial.view(*new_context_layer_shape)
622 context_layer = (1 - self.spatial_ratio) * context_layer + self.spatial_ratio * context_layer_spatial
624 hidden_states = self.dense(context_layer)
625 hidden_states = self.out_dropout(hidden_states)
626 hidden_states = self.LayerNorm(hidden_states + input_tensor)
627 return hidden_states
630class FeedForward(nn.Module):
631 """Point-wise feed-forward layer is implemented by two dense layers.
633 Args:
634 input_tensor (torch.Tensor): the input of the point-wise feed-forward layer
636 Returns:
637 hidden_states (torch.Tensor): the output of the point-wise feed-forward layer
639 """
641 def __init__(self, hidden_size, inner_size, hidden_dropout_prob, hidden_act, layer_norm_eps):
642 super().__init__()
643 self.dense_1 = nn.Linear(hidden_size, inner_size)
644 self.intermediate_act_fn = self.get_hidden_act(hidden_act)
646 self.dense_2 = nn.Linear(inner_size, hidden_size)
647 self.LayerNorm = nn.LayerNorm(hidden_size, eps=layer_norm_eps)
648 self.dropout = nn.Dropout(hidden_dropout_prob)
650 def get_hidden_act(self, act):
651 ACT2FN = {
652 "gelu": self.gelu,
653 "relu": fn.relu,
654 "swish": self.swish,
655 "tanh": torch.tanh,
656 "sigmoid": torch.sigmoid,
657 }
658 return ACT2FN[act]
660 def gelu(self, x):
661 """Implementation of the gelu activation function.
663 For information: OpenAI GPT's gelu is slightly different (and gives slightly different results)::
665 0.5 * x * (1 + torch.tanh(math.sqrt(2 / math.pi) * (x + 0.044715 * torch.pow(x, 3))))
667 Also see https://arxiv.org/abs/1606.08415
668 """
669 return x * 0.5 * (1.0 + torch.erf(x / math.sqrt(2.0)))
671 def swish(self, x):
672 return x * torch.sigmoid(x)
674 def forward(self, input_tensor):
675 hidden_states = self.dense_1(input_tensor)
676 hidden_states = self.intermediate_act_fn(hidden_states)
678 hidden_states = self.dense_2(hidden_states)
679 hidden_states = self.dropout(hidden_states)
680 hidden_states = self.LayerNorm(hidden_states + input_tensor)
682 return hidden_states
685class FEABlock(nn.Module):
686 """One transformer layer consists of a multi-head self-attention layer and a point-wise feed-forward layer.
688 Args:
689 hidden_states (torch.Tensor): the input of the multi-head self-attention sublayer
690 attention_mask (torch.Tensor): the attention mask for the multi-head self-attention sublayer
692 Returns:
693 feedforward_output (torch.Tensor): The output of the point-wise feed-forward sublayer,
694 is the output of the transformer layer.
696 """
698 def __init__(
699 self,
700 n_heads,
701 hidden_size,
702 intermediate_size,
703 hidden_dropout_prob,
704 attn_dropout_prob,
705 hidden_act,
706 layer_norm_eps,
707 n,
708 config,
709 ):
710 super().__init__()
711 self.hybrid_attention = HybridAttention(
712 n_heads,
713 hidden_size,
714 hidden_dropout_prob,
715 attn_dropout_prob,
716 layer_norm_eps,
717 n,
718 config,
719 )
720 self.feed_forward = FeedForward(
721 hidden_size,
722 intermediate_size,
723 hidden_dropout_prob,
724 hidden_act,
725 layer_norm_eps,
726 )
728 def forward(self, hidden_states, attention_mask):
729 attention_output = self.hybrid_attention(hidden_states, attention_mask)
730 feedforward_output = self.feed_forward(attention_output)
732 return feedforward_output
735class FEAEncoder(nn.Module):
736 r"""One TransformerEncoder consists of several TransformerLayers.
738 - n_layers(num): num of transformer layers in transformer encoder. Default: 2
739 - n_heads(num): num of attention heads for multi-head attention layer. Default: 2
740 - hidden_size(num): the input and output hidden size. Default: 64
741 - inner_size(num): the dimensionality in feed-forward layer. Default: 256
742 - hidden_dropout_prob(float): probability of an element to be zeroed. Default: 0.5
743 - attn_dropout_prob(float): probability of an attention score to be zeroed. Default: 0.5
744 - hidden_act(str): activation function in feed-forward layer. Default: 'gelu'
745 candidates: 'gelu', 'relu', 'swish', 'tanh', 'sigmoid'
746 - layer_norm_eps(float): a value added to the denominator for numerical stability. Default: 1e-12
748 """
750 def __init__(
751 self,
752 n_layers=2,
753 n_heads=2,
754 hidden_size=64,
755 inner_size=256,
756 hidden_dropout_prob=0.5,
757 attn_dropout_prob=0.5,
758 hidden_act="gelu",
759 layer_norm_eps=1e-12,
760 config=None,
761 ):
762 super().__init__()
763 self.n_layers = n_layers
764 self.layer = nn.ModuleList()
765 for n in range(self.n_layers):
766 self.layer_ramp = FEABlock(
767 n_heads,
768 hidden_size,
769 inner_size,
770 hidden_dropout_prob,
771 attn_dropout_prob,
772 hidden_act,
773 layer_norm_eps,
774 n,
775 config,
776 )
777 self.layer.append(self.layer_ramp)
779 def forward(self, hidden_states, attention_mask, output_all_encoded_layers=True):
780 """Args:
781 hidden_states (torch.Tensor): the input of the TransformerEncoder
782 attention_mask (torch.Tensor): the attention mask for the input hidden_states
783 output_all_encoded_layers (Bool): whether output all transformer layers' output
785 Returns:
786 all_encoder_layers (list): if output_all_encoded_layers is True, return a list consists of all transformer
787 layers' output, otherwise return a list only consists of the output of last transformer layer.
789 """
790 all_encoder_layers = []
792 for layer_module in self.layer:
793 hidden_states = layer_module(hidden_states, attention_mask)
794 if output_all_encoded_layers:
795 all_encoder_layers.append(hidden_states)
796 if not output_all_encoded_layers:
797 all_encoder_layers.append(hidden_states)
798 return all_encoder_layers