hanjian.thu123 commited on
Commit
217faf6
·
1 Parent(s): 0a1e9d5

[update] load device

Browse files
Files changed (2) hide show
  1. README.md +1 -1
  2. grn/utils_t2iv/load.py +4 -2
README.md CHANGED
@@ -6,7 +6,7 @@ emoji: 🚀
6
  colorFrom: red
7
  colorTo: yellow
8
  pinned: true
9
- short_description: Text-to-Image Demo for "Generative Refinement Networks"
10
  ---
11
  # GRN: Generative Refinement Networks
12
 
 
6
  colorFrom: red
7
  colorTo: yellow
8
  pinned: true
9
+ short_description: "Generative Refinement Networks"
10
  ---
11
  # GRN: Generative Refinement Networks
12
 
grn/utils_t2iv/load.py CHANGED
@@ -16,12 +16,14 @@ from timm.models import create_model
16
  def load_visual_tokenizer(args, device=None):
17
  if not device:
18
  device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
 
 
19
  vae = HBQ_Tokenizer(args=args, latent_channels=args.detail_scale_dim, encoder_out_type='feature_tanh')
20
  vae.eval()
21
- vae = vae.to('cuda')
22
  for param in vae.parameters():
23
  param.requires_grad = False
24
- state_dict = torch.load(args.vae_path, map_location='cuda')
25
  if 'ema' in state_dict:
26
  print(f'Load ema vae weights')
27
  state_dict = state_dict['ema']
 
16
  def load_visual_tokenizer(args, device=None):
17
  if not device:
18
  device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
19
+ elif isinstance(device, str):
20
+ device = torch.device(device)
21
  vae = HBQ_Tokenizer(args=args, latent_channels=args.detail_scale_dim, encoder_out_type='feature_tanh')
22
  vae.eval()
23
+ vae = vae.to(device)
24
  for param in vae.parameters():
25
  param.requires_grad = False
26
+ state_dict = torch.load(args.vae_path, map_location=device)
27
  if 'ema' in state_dict:
28
  print(f'Load ema vae weights')
29
  state_dict = state_dict['ema']