Spaces:
Running on Zero
Running on Zero
hanjian.thu123 commited on
Commit ·
217faf6
1
Parent(s): 0a1e9d5
[update] load device
Browse files- README.md +1 -1
- 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:
|
| 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(
|
| 22 |
for param in vae.parameters():
|
| 23 |
param.requires_grad = False
|
| 24 |
-
state_dict = torch.load(args.vae_path, map_location=
|
| 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']
|