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

1# @Time : 2023/10/27 

2# @Author : Kesha Ou 

3# @Email : keishaou@gmail.com 

4 

5r"""FEARec 

6################################################ 

7 

8Reference: 

9 Xinyu Du et al. "Frequency Enhanced Hybrid Attention Network for Sequential Recommendation." 

10 In SIGIR 2023. 

11 

12Reference code: 

13 https://github.com/sudaada/FEARec 

14 

15""" 

16 

17import math 

18import random 

19 

20import numpy as np 

21import torch 

22import torch.nn.functional as F 

23import torch.nn.functional as fn 

24from torch import nn 

25 

26from hopwise.data.interaction import Interaction 

27from hopwise.model.abstract_recommender import SequentialRecommender 

28from hopwise.model.loss import BPRLoss 

29 

30 

31class FEARec(SequentialRecommender): 

32 def __init__(self, config, dataset): 

33 super().__init__(config, dataset) 

34 

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"] 

46 

47 self.lmd = config["lmd"] 

48 self.lmd_sem = config["lmd_sem"] 

49 

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) 

53 

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 ) 

68 

69 self.LayerNorm = nn.LayerNorm(self.hidden_size, eps=self.layer_norm_eps) 

70 self.dropout = nn.Dropout(self.hidden_dropout_prob) 

71 

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']!") 

78 

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() 

88 

89 # parameters initialization 

90 self.apply(self._init_weights) 

91 

92 def get_same_item_index(self, dataset): 

93 same_target_index = {} 

94 target_item = dataset.inter_feat[self.ITEM_ID].numpy() 

95 

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 

99 

100 return same_target_index 

101 

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_() 

114 

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 

124 

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) 

135 

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 

140 

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 

149 

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) 

154 

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) 

159 

160 extended_attention_mask = self.get_attention_mask(item_seq) 

161 # extended_attention_mask = self.get_bi_attention_mask(item_seq) 

162 

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) 

166 

167 return output # [B H] 

168 

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() 

173 

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() 

179 

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] 

200 

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 

205 

206 sem_pos_lengths = torch.stack(sem_pos_lengths).to(self.device) 

207 sem_pos_seqs = torch.stack(sem_pos_seqs).to(self.device) 

208 

209 interaction.update(Interaction({"sem_aug": sem_pos_seqs, "sem_aug_lengths": sem_pos_lengths})) 

210 

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) 

226 

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 ) 

237 

238 loss += self.lmd * self.aug_nce_fct(nce_logits, nce_labels) 

239 

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) 

247 

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 ) 

255 

256 loss += self.lmd_sem * self.aug_nce_fct(sem_nce_logits, sem_nce_labels) 

257 

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 ) 

264 

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 ) 

273 

274 loss += self.lmd_sem * self.aug_nce_fct(sem_nce_logits, sem_nce_labels) 

275 

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() 

287 

288 return loss 

289 

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 

298 

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 

304 

305 z = torch.cat((z_i, z_j), dim=0) 

306 

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 

311 

312 sim_i_j = torch.diag(sim, batch_size) 

313 sim_j_i = torch.diag(sim, -batch_size) 

314 

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) 

321 

322 labels = torch.zeros(N).to(positive_samples.device).long() 

323 logits = torch.cat((positive_samples, negative_samples), dim=1) 

324 return logits, labels 

325 

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 

331 

332 z = torch.cat((z_i, z_j), dim=0) 

333 

334 # pairwise l2 distace 

335 sim = torch.cdist(z, z, p=2) 

336 

337 sim_i_j = torch.diag(sim, batch_size) 

338 sim_j_i = torch.diag(sim, -batch_size) 

339 

340 positive_samples = torch.cat((sim_i_j, sim_j_i), dim=0).reshape(N, 1) 

341 alignment = positive_samples.mean() 

342 

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()) 

349 

350 return alignment, uniformity 

351 

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 

360 

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 

368 

369 

370class HybridAttention(nn.Module): 

371 """Hybrid Attention layer: combine time domain self-attention layer and frequency domain attention layer. 

372 

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 

376 

377 Returns: 

378 hidden_states (torch.Tensor): the output of the multi-head Hybrid Attention layer 

379 

380 """ 

381 

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 ) 

397 

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) 

428 

429 if self.filter_mixer == "G": 

430 self.w = self.global_ratio 

431 self.s = self.slide_step 

432 

433 if self.filter_mixer == "L": 

434 self.w = self.local_ratio 

435 self.s = self.filter_size 

436 

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) 

439 

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)) 

453 

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}") 

457 

458 if self.config["dual_domain"]: 

459 self.spatial_ratio = self.config["spatial_ratio"] 

460 

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 

469 

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 

493 

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 

529 

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) 

534 

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) 

538 

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, :, :] 

550 

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) 

554 

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] 

558 

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) 

562 

563 if self.config["use_filter"]: 

564 # filter 

565 weight = torch.view_as_complex(self.complex_weight) 

566 res = res * weight 

567 

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 

570 

571 corr = torch.fft.irfft(box_res, dim=-1) 

572 

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) 

578 

579 new_context_layer_shape = V.size()[:-2] + (self.all_head_size,) 

580 context_layer = V.view(*new_context_layer_shape) 

581 

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 

589 

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 

595 

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 

603 

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) 

607 

608 queries = queries.permute(0, 1, 3, 2) 

609 keys = keys.permute(0, 1, 3, 2) 

610 values = values.permute(0, 1, 3, 2) 

611 

612 attention_scores = torch.matmul(queries, keys.transpose(-1, -2)) 

613 attention_scores = attention_scores / math.sqrt(self.attention_head_size) 

614 

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 

623 

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 

628 

629 

630class FeedForward(nn.Module): 

631 """Point-wise feed-forward layer is implemented by two dense layers. 

632 

633 Args: 

634 input_tensor (torch.Tensor): the input of the point-wise feed-forward layer 

635 

636 Returns: 

637 hidden_states (torch.Tensor): the output of the point-wise feed-forward layer 

638 

639 """ 

640 

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) 

645 

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) 

649 

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] 

659 

660 def gelu(self, x): 

661 """Implementation of the gelu activation function. 

662 

663 For information: OpenAI GPT's gelu is slightly different (and gives slightly different results):: 

664 

665 0.5 * x * (1 + torch.tanh(math.sqrt(2 / math.pi) * (x + 0.044715 * torch.pow(x, 3)))) 

666 

667 Also see https://arxiv.org/abs/1606.08415 

668 """ 

669 return x * 0.5 * (1.0 + torch.erf(x / math.sqrt(2.0))) 

670 

671 def swish(self, x): 

672 return x * torch.sigmoid(x) 

673 

674 def forward(self, input_tensor): 

675 hidden_states = self.dense_1(input_tensor) 

676 hidden_states = self.intermediate_act_fn(hidden_states) 

677 

678 hidden_states = self.dense_2(hidden_states) 

679 hidden_states = self.dropout(hidden_states) 

680 hidden_states = self.LayerNorm(hidden_states + input_tensor) 

681 

682 return hidden_states 

683 

684 

685class FEABlock(nn.Module): 

686 """One transformer layer consists of a multi-head self-attention layer and a point-wise feed-forward layer. 

687 

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 

691 

692 Returns: 

693 feedforward_output (torch.Tensor): The output of the point-wise feed-forward sublayer, 

694 is the output of the transformer layer. 

695 

696 """ 

697 

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 ) 

727 

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) 

731 

732 return feedforward_output 

733 

734 

735class FEAEncoder(nn.Module): 

736 r"""One TransformerEncoder consists of several TransformerLayers. 

737 

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 

747 

748 """ 

749 

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) 

778 

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 

784 

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. 

788 

789 """ 

790 all_encoder_layers = [] 

791 

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