MCPcopy Create free account
hub / github.com/LTH14/rcg / forward_encoder

Method forward_encoder

pixel_generator/mage/models_mage.py:337–407  ·  view source on GitHub ↗
(self, x, rep, class_label)

Source from the content-addressed store, hash-verified

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

Callers 1

forwardMethod · 0.95

Calls 2

printFunction · 0.85
encodeMethod · 0.45

Tested by

no test coverage detected