| 335 | nn.init.constant_(m.weight, 1.0) |
| 336 | |
| 337 | def forward_encoder(self, x, rep, class_label): |
| 338 | # tokenization |
| 339 | with torch.no_grad(): |
| 340 | z_q, _, token_tuple = self.vqgan.encode(x) |
| 341 | |
| 342 | _, _, token_indices = token_tuple |
| 343 | token_indices = token_indices.reshape(z_q.size(0), -1) |
| 344 | gt_indices = token_indices.clone().detach().long() |
| 345 | |
| 346 | # masking |
| 347 | bsz, seq_len = token_indices.size() |
| 348 | mask_ratio_min = self.mask_ratio_min |
| 349 | mask_rate = self.mask_ratio_generator.rvs(1)[0] |
| 350 | |
| 351 | num_dropped_tokens = int(np.ceil(seq_len * mask_ratio_min)) |
| 352 | num_masked_tokens = int(np.ceil(seq_len * mask_rate)) |
| 353 | |
| 354 | # it is possible that two elements of the noise is the same, so do a while loop to avoid it |
| 355 | while True: |
| 356 | noise = torch.rand(bsz, seq_len, device=x.device) # noise in [0, 1] |
| 357 | sorted_noise, _ = torch.sort(noise, dim=1) # ascend: small is remove, large is keep |
| 358 | cutoff_drop = sorted_noise[:, num_dropped_tokens-1:num_dropped_tokens] |
| 359 | cutoff_mask = sorted_noise[:, num_masked_tokens-1:num_masked_tokens] |
| 360 | token_drop_mask = (noise <= cutoff_drop).float() |
| 361 | token_all_mask = (noise <= cutoff_mask).float() |
| 362 | if token_drop_mask.sum() == bsz*num_dropped_tokens and token_all_mask.sum() == bsz*num_masked_tokens: |
| 363 | break |
| 364 | else: |
| 365 | print("Rerandom the noise!") |
| 366 | # print(mask_rate, num_dropped_tokens, num_masked_tokens, token_drop_mask.sum(dim=1), token_all_mask.sum(dim=1)) |
| 367 | token_indices[token_all_mask.nonzero(as_tuple=True)] = self.mask_token_label |
| 368 | # print("Masekd num token:", torch.sum(token_indices == self.mask_token_label, dim=1)) |
| 369 | |
| 370 | # concate class token |
| 371 | token_indices = torch.cat([torch.zeros(token_indices.size(0), 1).cuda(device=token_indices.device), token_indices], dim=1) |
| 372 | token_indices[:, 0] = self.fake_class_label |
| 373 | token_drop_mask = torch.cat([torch.zeros(token_indices.size(0), 1).cuda(), token_drop_mask], dim=1) |
| 374 | token_all_mask = torch.cat([torch.zeros(token_indices.size(0), 1).cuda(), token_all_mask], dim=1) |
| 375 | token_indices = token_indices.long() |
| 376 | # bert embedding |
| 377 | input_embeddings = self.token_emb(token_indices) |
| 378 | # print("Input embedding shape:", input_embeddings.shape) |
| 379 | bsz, seq_len, emb_dim = input_embeddings.shape |
| 380 | |
| 381 | # dropping |
| 382 | token_keep_mask = 1 - token_drop_mask |
| 383 | input_embeddings_after_drop = input_embeddings[token_keep_mask.nonzero(as_tuple=True)].reshape(bsz, -1, emb_dim) |
| 384 | # print("Input embedding after drop shape:", input_embeddings_after_drop.shape) |
| 385 | |
| 386 | # replace fake class token with rep |
| 387 | if self.use_rep: |
| 388 | # cfg by masking representation |
| 389 | drop_rep_mask = torch.rand(bsz) < self.rep_drop_prob |
| 390 | drop_rep_mask = drop_rep_mask.unsqueeze(-1).cuda().float() |
| 391 | rep = drop_rep_mask * self.fake_latent + (1 - drop_rep_mask) * rep |
| 392 | |
| 393 | rep = self.latent_prior_proj(rep) |
| 394 | input_embeddings_after_drop[:, 0] = rep |