Coverage for hopwise/model/general_recommender/diffrec.py: 75%

309 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-09-30 13:25 +0000

1# @Time : 2023/10/6 

2# @Author : Enze Liu 

3# @Email : enzeeliu@foxmail.com 

4 

5r"""DiffRec 

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

7Reference: 

8 Wenjie Wang et al. "Diffusion Recommender Model." in SIGIR 2023. 

9 

10Reference code: 

11 https://github.com/YiyanXu/DiffRec 

12""" 

13 

14import copy 

15import enum 

16import math 

17 

18import numpy as np 

19import torch 

20import torch.nn.functional as F 

21from torch import nn 

22 

23from hopwise.model.abstract_recommender import AutoEncoderMixin, GeneralRecommender 

24from hopwise.model.init import xavier_normal_initialization 

25from hopwise.model.layers import MLPLayers 

26from hopwise.utils import InputType 

27 

28 

29class ModelMeanType(enum.Enum): 

30 START_X = enum.auto() # the model predicts x_0 

31 EPSILON = enum.auto() # the model predicts epsilon 

32 

33 

34class DNN(nn.Module): 

35 r"""A deep neural network for the reverse diffusion preocess.""" 

36 

37 def __init__( 

38 self, 

39 dims: list, 

40 emb_size: int, 

41 time_type="cat", 

42 act_func="tanh", 

43 norm=False, 

44 dropout=0.5, 

45 ): 

46 super().__init__() 

47 self.dims = dims 

48 self.time_type = time_type 

49 self.time_emb_dim = emb_size 

50 self.norm = norm 

51 

52 self.emb_layer = nn.Linear(self.time_emb_dim, self.time_emb_dim) 

53 

54 if self.time_type == "cat": 

55 # Concatenate timestep embedding with input 

56 self.dims[0] += self.time_emb_dim 

57 else: 

58 raise ValueError("Unimplemented timestep embedding type %s" % self.time_type) 

59 

60 self.mlp_layers = MLPLayers(layers=self.dims, dropout=0, activation=act_func, last_activation=False) 

61 self.drop = nn.Dropout(dropout) 

62 

63 self.apply(xavier_normal_initialization) 

64 

65 def forward(self, x, timesteps): 

66 time_emb = timestep_embedding(timesteps, self.time_emb_dim).to(x.device) 

67 emb = self.emb_layer(time_emb) 

68 if self.norm: 

69 x = F.normalize(x) 

70 x = self.drop(x) 

71 h = torch.cat([x, emb], dim=-1) 

72 h = self.mlp_layers(h) 

73 return h 

74 

75 

76class DiffRec(GeneralRecommender, AutoEncoderMixin): 

77 r"""DiffRec is a generative recommender model which infers users' interaction probabilities in a denoising manner. 

78 Note that DiffRec simultaneously ranks all items for each user. 

79 We implement the the DiffRec model with only user dataloader. 

80 """ 

81 

82 input_type = InputType.USERWISE 

83 

84 def __init__(self, config, dataset): 

85 super().__init__(config, dataset) 

86 

87 if config["mean_type"] == "x0": 

88 self.mean_type = ModelMeanType.START_X 

89 elif config["mean_type"] == "eps": 

90 self.mean_type = ModelMeanType.EPSILON 

91 else: 

92 raise ValueError("Unimplemented mean type %s" % config["mean_type"]) 

93 self.time_aware = config["time-aware"] 

94 self.w_max = config["w_max"] 

95 self.w_min = config["w_min"] 

96 self.build_histroy_items(dataset) 

97 

98 self.noise_schedule = config["noise_schedule"] 

99 self.noise_scale = config["noise_scale"] 

100 self.noise_min = config["noise_min"] 

101 self.noise_max = config["noise_max"] 

102 self.steps = config["steps"] 

103 self.beta_fixed = config["beta_fixed"] 

104 self.emb_size = config["embedding_size"] 

105 self.norm = config["norm"] # True or False 

106 self.reweight = config["reweight"] # reweight the loss for different timesteps 

107 if self.noise_scale == 0.0: 

108 self.reweight = False 

109 self.sampling_noise = config["sampling_noise"] # whether sample noise during predict 

110 self.sampling_steps = config["sampling_steps"] 

111 self.mlp_act_func = config["mlp_act_func"] 

112 assert self.sampling_steps <= self.steps, "Too much steps in inference." 

113 

114 self.history_num_per_term = config["history_num_per_term"] 

115 self.Lt_history = torch.zeros(self.steps, self.history_num_per_term, dtype=torch.float64).to(self.device) 

116 self.Lt_count = torch.zeros(self.steps, dtype=int).to(self.device) 

117 

118 dims = [self.n_items] + config["dims_dnn"] + [self.n_items] 

119 

120 self.mlp = DNN( 

121 dims=dims, 

122 emb_size=self.emb_size, 

123 time_type="cat", 

124 norm=self.norm, 

125 act_func=self.mlp_act_func, 

126 ).to(self.device) 

127 

128 if self.noise_scale != 0.0: 

129 self.betas = torch.tensor(self.get_betas(), dtype=torch.float64).to(self.device) 

130 if self.beta_fixed: 

131 self.betas[0] = 0.00001 # Deep Unsupervised Learning using Noneequilibrium Thermodynamics 2.4.1 

132 # The variance \beta_1 of the first step is fixed to a small constant to prevent overfitting. 

133 assert len(self.betas.shape) == 1, "betas must be 1-D" 

134 assert len(self.betas) == self.steps, "num of betas must equal to diffusion steps" 

135 assert (self.betas > 0).all() and (self.betas <= 1).all(), "betas out of range" 

136 

137 self.calculate_for_diffusion() 

138 

139 def build_histroy_items(self, dataset): 

140 r"""Add time-aware reweighting to the original user-item interaction matrix when config['time-aware'] is True.""" # noqa: E501 

141 if not self.time_aware: 

142 super().build_histroy_items(dataset) 

143 else: 

144 inter_feat = copy.deepcopy(dataset.inter_feat) 

145 inter_feat.sort(dataset.time_field) 

146 user_ids, item_ids = ( 

147 inter_feat[dataset.uid_field].numpy(), 

148 inter_feat[dataset.iid_field].numpy(), 

149 ) 

150 

151 w_max = self.w_max 

152 w_min = self.w_min 

153 values = np.zeros(len(inter_feat)) 

154 

155 row_num = dataset.user_num 

156 row_ids, col_ids = user_ids, item_ids 

157 

158 for uid in range(1, row_num + 1): 

159 uindex = np.argwhere(user_ids == uid).flatten() 

160 int_num = len(uindex) 

161 weight = np.linspace(w_min, w_max, int_num) 

162 values[uindex] = weight 

163 

164 history_len = np.zeros(row_num, dtype=np.int64) 

165 for row_id in row_ids: 

166 history_len[row_id] += 1 

167 

168 max_inter_num = np.max(history_len) 

169 col_num = max_inter_num 

170 

171 history_matrix = np.zeros((row_num, col_num), dtype=np.int64) 

172 history_value = np.zeros((row_num, col_num)) 

173 history_len[:] = 0 

174 

175 for row_id, value, col_id in zip(row_ids, values, col_ids): 

176 if history_len[row_id] >= col_num: 

177 continue 

178 history_matrix[row_id, history_len[row_id]] = col_id 

179 history_value[row_id, history_len[row_id]] = value 

180 history_len[row_id] += 1 

181 

182 self.history_item_id = torch.LongTensor(history_matrix) 

183 self.history_item_value = torch.FloatTensor(history_value) 

184 self.history_item_id = self.history_item_id.to(self.device) 

185 self.history_item_value = self.history_item_value.to(self.device) 

186 

187 def get_betas(self): 

188 r"""Given the schedule name, create the betas for the diffusion process.""" 

189 if self.noise_schedule in ("linear", "linear-var"): 

190 start = self.noise_scale * self.noise_min 

191 end = self.noise_scale * self.noise_max 

192 if self.noise_schedule == "linear": 

193 return np.linspace(start, end, self.steps, dtype=np.float64) 

194 else: 

195 return betas_from_linear_variance(self.steps, np.linspace(start, end, self.steps, dtype=np.float64)) 

196 elif self.noise_schedule == "cosine": 

197 return betas_for_alpha_bar(self.steps, lambda t: math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2) 

198 # Deep Unsupervised Learning using Noneequilibrium Thermodynamics 2.4.1 

199 elif self.noise_schedule == "binomial": 

200 ts = np.arange(self.steps) 

201 betas = [1 / (self.steps - t + 1) for t in ts] 

202 return betas 

203 else: 

204 raise NotImplementedError(f"unknown beta schedule: {self.noise_schedule}!") 

205 

206 def calculate_for_diffusion(self): 

207 r"""Calculate the coefficients for the diffusion process.""" 

208 alphas = 1.0 - self.betas 

209 # [alpha_{1}, ..., alpha_{1}*...*alpha_{T}] shape (steps,) 

210 self.alphas_cumprod = torch.cumprod(alphas, axis=0).to(self.device) 

211 # alpha_{t-1} 

212 self.alphas_cumprod_prev = torch.cat([torch.tensor([1.0]).to(self.device), self.alphas_cumprod[:-1]]).to( 

213 self.device 

214 ) 

215 # alpha_{t+1} 

216 self.alphas_cumprod_next = torch.cat([self.alphas_cumprod[1:], torch.tensor([0.0]).to(self.device)]).to( 

217 self.device 

218 ) 

219 assert self.alphas_cumprod_prev.shape == (self.steps,) 

220 

221 self.sqrt_alphas_cumprod = torch.sqrt(self.alphas_cumprod) 

222 self.sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - self.alphas_cumprod) 

223 self.log_one_minus_alphas_cumprod = torch.log(1.0 - self.alphas_cumprod) 

224 self.sqrt_recip_alphas_cumprod = torch.sqrt(1.0 / self.alphas_cumprod) 

225 self.sqrt_recipm1_alphas_cumprod = torch.sqrt(1.0 / self.alphas_cumprod - 1) 

226 

227 self.posterior_variance = self.betas * (1.0 - self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod) 

228 

229 self.posterior_log_variance_clipped = torch.log( 

230 torch.cat([self.posterior_variance[1].unsqueeze(0), self.posterior_variance[1:]]) 

231 ) 

232 # Eq.10 coef for x_theta 

233 self.posterior_mean_coef1 = self.betas * torch.sqrt(self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod) 

234 # Eq.10 coef for x_t 

235 self.posterior_mean_coef2 = (1.0 - self.alphas_cumprod_prev) * torch.sqrt(alphas) / (1.0 - self.alphas_cumprod) 

236 

237 def p_sample(self, x_start): 

238 r"""Generate users' interaction probabilities in a denoising manner. 

239 

240 Args: 

241 x_start (torch.FloatTensor): the input tensor that contains user's history interaction matrix, 

242 for DiffRec shape: [batch_size, n_items] 

243 for LDiffRec shape: [batch_size, hidden_size] 

244 

245 Returns: 

246 torch.FloatTensor: the interaction probabilities, 

247 for DiffRec shape: [batch_size, n_items] 

248 for LDiffRec shape: [batch_size, hidden_size] 

249 """ 

250 steps = self.sampling_steps 

251 if steps == 0: 

252 x_t = x_start 

253 else: 

254 t = torch.tensor([steps - 1] * x_start.shape[0]).to(x_start.device) 

255 x_t = self.q_sample(x_start, t) 

256 

257 indices = list(range(self.steps))[::-1] 

258 

259 if self.noise_scale == 0.0: 

260 for i in indices: 

261 t = torch.tensor([i] * x_t.shape[0]).to(x_start.device) 

262 x_t = self.mlp(x_t, t) 

263 return x_t 

264 

265 for i in indices: 

266 t = torch.tensor([i] * x_t.shape[0]).to(x_start.device) 

267 out = self.p_mean_variance(x_t, t) 

268 if self.sampling_noise: 

269 noise = torch.randn_like(x_t) 

270 nonzero_mask = (t != 0).float().view(-1, *([1] * (len(x_t.shape) - 1))) # no noise when t == 0 

271 x_t = out["mean"] + nonzero_mask * torch.exp(0.5 * out["log_variance"]) * noise 

272 else: 

273 x_t = out["mean"] 

274 return x_t 

275 

276 def full_sort_predict(self, interaction): 

277 user = interaction[self.USER_ID] 

278 x_start = self.get_rating_matrix(user) 

279 scores = self.p_sample(x_start) 

280 return scores 

281 

282 def predict(self, interaction): 

283 item = interaction[self.ITEM_ID] 

284 x_t = self.full_sort_predict(interaction) 

285 scores = x_t[torch.arange(len(item)).to(self.device), item] 

286 return scores 

287 

288 def calculate_loss(self, interaction): 

289 user = interaction[self.USER_ID] 

290 x_start = self.get_rating_matrix(user) 

291 

292 batch_size, device = x_start.size(0), x_start.device 

293 ts, pt = self.sample_timesteps(batch_size, device, "importance") 

294 noise = torch.randn_like(x_start) 

295 if self.noise_scale != 0.0: 

296 x_t = self.q_sample(x_start, ts, noise) 

297 else: 

298 x_t = x_start 

299 

300 model_output = self.mlp(x_t, ts) 

301 target = { 

302 ModelMeanType.START_X: x_start, 

303 ModelMeanType.EPSILON: noise, 

304 }[self.mean_type] 

305 

306 assert model_output.shape == target.shape == x_start.shape 

307 

308 mse = mean_flat((target - model_output) ** 2) 

309 

310 reloss = self.reweight_loss(x_start, x_t, mse, ts, target, model_output, device) 

311 self.update_Lt_history(ts, reloss) 

312 

313 # importance sampling 

314 reloss /= pt 

315 mean_loss = reloss.mean() 

316 return mean_loss 

317 

318 def reweight_loss(self, x_start, x_t, mse, ts, target, model_output, device): 

319 if self.reweight: 

320 if self.mean_type == ModelMeanType.START_X: 

321 # Eq.11 

322 weight = self.SNR(ts - 1) - self.SNR(ts) 

323 # Eq.12 

324 weight = torch.where((ts == 0), 1.0, weight) 

325 loss = mse 

326 elif self.mean_type == ModelMeanType.EPSILON: 

327 weight = (1 - self.alphas_cumprod[ts]) / ( 

328 (1 - self.alphas_cumprod_prev[ts]) ** 2 * (1 - self.betas[ts]) 

329 ) 

330 weight = torch.where((ts == 0), 1.0, weight) 

331 likelihood = mean_flat((x_start - self._predict_xstart_from_eps(x_t, ts, model_output)) ** 2 / 2.0) 

332 loss = torch.where((ts == 0), likelihood, mse) 

333 else: 

334 weight = torch.tensor([1.0] * len(target)).to(device) 

335 loss = mse 

336 reloss = weight * loss 

337 return reloss 

338 

339 def update_Lt_history(self, ts, reloss): 

340 # update Lt_history & Lt_count 

341 for t, loss in zip(ts, reloss): 

342 if self.Lt_count[t] == self.history_num_per_term: 

343 Lt_history_old = self.Lt_history.clone() 

344 self.Lt_history[t, :-1] = Lt_history_old[t, 1:] 

345 self.Lt_history[t, -1] = loss.detach() 

346 else: 

347 try: 

348 self.Lt_history[t, self.Lt_count[t]] = loss.detach() 

349 self.Lt_count[t] += 1 

350 except Exception: 

351 print(t) 

352 print(self.Lt_count[t]) 

353 print(loss) 

354 raise ValueError 

355 

356 def sample_timesteps(self, batch_size, device, method="uniform", uniform_prob=0.001): 

357 if method == "importance": # importance sampling 

358 if not (self.Lt_count == self.history_num_per_term).all(): 

359 return self.sample_timesteps(batch_size, device, method="uniform") 

360 

361 Lt_sqrt = torch.sqrt(torch.mean(self.Lt_history**2, axis=-1)) 

362 pt_all = Lt_sqrt / torch.sum(Lt_sqrt) 

363 pt_all *= 1 - uniform_prob 

364 pt_all += uniform_prob / len(pt_all) # ensure the least prob > uniform_prob 

365 

366 assert pt_all.sum(-1) - 1.0 < 1e-5 # noqa: PLR2004 

367 

368 t = torch.multinomial(pt_all, num_samples=batch_size, replacement=True) 

369 pt = pt_all.gather(dim=0, index=t) * len(pt_all) 

370 

371 return t, pt 

372 

373 elif method == "uniform": # uniform sampling 

374 t = torch.randint(0, self.steps, (batch_size,), device=device).long() 

375 pt = torch.ones_like(t).float() 

376 

377 return t, pt 

378 

379 else: 

380 raise ValueError 

381 

382 def q_sample(self, x_start, t, noise=None): 

383 if noise is None: 

384 noise = torch.randn_like(x_start) 

385 assert noise.shape == x_start.shape 

386 return ( 

387 self._extract_into_tensor(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start 

388 + self._extract_into_tensor(self.sqrt_one_minus_alphas_cumprod, t, x_start.shape) * noise 

389 ) 

390 

391 def q_posterior_mean_variance(self, x_start, x_t, t): 

392 r"""Compute the mean and variance of the diffusion posterior: 

393 q(x_{t-1} | x_t, x_0) 

394 """ 

395 assert x_start.shape == x_t.shape 

396 posterior_mean = ( 

397 self._extract_into_tensor(self.posterior_mean_coef1, t, x_t.shape) * x_start 

398 + self._extract_into_tensor(self.posterior_mean_coef2, t, x_t.shape) * x_t 

399 ) 

400 posterior_variance = self._extract_into_tensor(self.posterior_variance, t, x_t.shape) 

401 posterior_log_variance_clipped = self._extract_into_tensor(self.posterior_log_variance_clipped, t, x_t.shape) 

402 assert ( 

403 posterior_mean.shape[0] 

404 == posterior_variance.shape[0] 

405 == posterior_log_variance_clipped.shape[0] 

406 == x_start.shape[0] 

407 ) 

408 return posterior_mean, posterior_variance, posterior_log_variance_clipped 

409 

410 def p_mean_variance(self, x, t): 

411 r"""Apply the model to get p(x_{t-1} | x_t), as well as a prediction of 

412 the initial x, x_0. 

413 """ 

414 B, C = x.shape[:2] 

415 assert t.shape == (B,) 

416 model_output = self.mlp(x, t) 

417 

418 model_variance = self.posterior_variance 

419 model_log_variance = self.posterior_log_variance_clipped 

420 

421 model_variance = self._extract_into_tensor(model_variance, t, x.shape) 

422 model_log_variance = self._extract_into_tensor(model_log_variance, t, x.shape) 

423 

424 if self.mean_type == ModelMeanType.START_X: 

425 pred_xstart = model_output 

426 elif self.mean_type == ModelMeanType.EPSILON: 

427 pred_xstart = self._predict_xstart_from_eps(x, t, eps=model_output) 

428 else: 

429 raise NotImplementedError(self.mean_type) 

430 

431 model_mean, _, _ = self.q_posterior_mean_variance(x_start=pred_xstart, x_t=x, t=t) 

432 

433 assert model_mean.shape == model_log_variance.shape == pred_xstart.shape == x.shape 

434 

435 return { 

436 "mean": model_mean, 

437 "variance": model_variance, 

438 "log_variance": model_log_variance, 

439 "pred_xstart": pred_xstart, 

440 } 

441 

442 def _predict_xstart_from_eps(self, x_t, t, eps): 

443 assert x_t.shape == eps.shape 

444 return ( 

445 self._extract_into_tensor(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t 

446 - self._extract_into_tensor(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape) * eps 

447 ) 

448 

449 def SNR(self, t): 

450 r"""Compute the signal-to-noise ratio for a single timestep.""" 

451 self.alphas_cumprod = self.alphas_cumprod.to(t.device) 

452 return self.alphas_cumprod[t] / (1 - self.alphas_cumprod[t]) 

453 

454 def _extract_into_tensor(self, arr, timesteps, broadcast_shape): 

455 r"""Extract values from a 1-D torch tensor for a batch of indices. 

456 

457 Args: 

458 arr (torch.Tensor): the 1-D torch tensor. 

459 timesteps (torch.Tensor): a tensor of indices into the array to extract. 

460 broadcast_shape (torch.Size): a larger shape of K dimensions with the batch 

461 dimension equal to the length of timesteps. 

462 

463 Returns: 

464 torch.Tensor: a tensor of shape [batch_size, 1, ...] where the shape has K dims. 

465 """ 

466 # res = torch.from_numpy(arr).to(device=timesteps.device)[timesteps].float() 

467 arr = arr.to(timesteps.device) 

468 res = arr[timesteps].float() 

469 while len(res.shape) < len(broadcast_shape): 

470 res = res[..., None] 

471 return res.expand(broadcast_shape) 

472 

473 

474def betas_from_linear_variance(steps, variance, max_beta=0.999): 

475 alpha_bar = 1 - variance 

476 betas = [] 

477 betas.append(1 - alpha_bar[0]) 

478 for i in range(1, steps): 

479 betas.append(min(1 - alpha_bar[i] / alpha_bar[i - 1], max_beta)) 

480 return np.array(betas) 

481 

482 

483def betas_for_alpha_bar(num_diffusion_timesteps, alpha_bar, max_beta=0.999): 

484 r"""Create a beta schedule that discretizes the given alpha_t_bar function, 

485 which defines the cumulative product of (1-beta) over time from t = [0,1]. 

486 

487 Args: 

488 num_diffusion_timesteps (int): the number of betas to produce. 

489 alpha_bar (Callable): a lambda that takes an argument t from 0 to 1 and 

490 produces the cumulative product of (1-beta) up to that 

491 part of the diffusion process. 

492 max_beta (int): the maximum beta to use; use values lower than 1 to 

493 prevent singularities. 

494 

495 Returns: 

496 np.ndarray: a 1-D array of beta values. 

497 """ 

498 betas = [] 

499 for i in range(num_diffusion_timesteps): 

500 t1 = i / num_diffusion_timesteps 

501 t2 = (i + 1) / num_diffusion_timesteps 

502 betas.append(min(1 - alpha_bar(t2) / alpha_bar(t1), max_beta)) 

503 return np.array(betas) 

504 

505 

506def normal_kl(mean1, logvar1, mean2, logvar2): 

507 r"""Compute the KL divergence between two gaussians. 

508 

509 Shapes are automatically broadcasted, so batches can be compared to 

510 scalars, among other use cases. 

511 """ 

512 tensor = None 

513 for obj in (mean1, logvar1, mean2, logvar2): 

514 if isinstance(obj, torch.Tensor): 

515 tensor = obj 

516 break 

517 assert tensor is not None, "at least one argument must be a Tensor" 

518 

519 # Force variances to be Tensors. Broadcasting helps convert scalars to 

520 # Tensors, but it does not work for torch.exp(). 

521 logvar1, logvar2 = (x if isinstance(x, torch.Tensor) else torch.tensor(x).to(tensor) for x in (logvar1, logvar2)) 

522 

523 return 0.5 * ( 

524 -1.0 + logvar2 - logvar1 + torch.exp(logvar1 - logvar2) + ((mean1 - mean2) ** 2) * torch.exp(-logvar2) 

525 ) 

526 

527 

528def mean_flat(tensor): 

529 r"""Take the mean over all non-batch dimensions.""" 

530 return tensor.mean(dim=list(range(1, len(tensor.shape)))) 

531 

532 

533def timestep_embedding(timesteps, dim, max_period=10000): 

534 r"""Create sinusoidal timestep embeddings. 

535 

536 :param timesteps: a 1-D Tensor of N indices, one per batch element. 

537 These may be fractional. (N,) 

538 :param dim: the dimension of the output. 

539 :param max_period: controls the minimum frequency of the embeddings. 

540 :return: an [N x dim] Tensor of positional embeddings. 

541 """ 

542 half = dim // 2 

543 freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half).to( 

544 timesteps.device 

545 ) # shape (dim//2,) 

546 args = timesteps[:, None].float() * freqs[None] # (N, dim//2) 

547 embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1) # (N, (dim//2)*2) 

548 if dim % 2: 

549 # zero pad in the last dimension to ensure shape (N, dim) 

550 embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1) 

551 return embedding