Team Ai
Apppublic

radames/Text2Human-API

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
patch170 linesDownload Raw Back to root
1diff --git a/models/hierarchy_inference_model.py b/models/hierarchy_inference_model.py2index 3116307..5de661d 1006443--- a/models/hierarchy_inference_model.py4+++ b/models/hierarchy_inference_model.py5@@ -21,7 +21,7 @@ class VQGANTextureAwareSpatialHierarchyInferenceModel():6 7     def __init__(self, opt):8         self.opt = opt9-        self.device = torch.device('cuda')10+        self.device = torch.device(opt['device'])11         self.is_train = opt['is_train']12 13         self.top_encoder = Encoder(14diff --git a/models/hierarchy_vqgan_model.py b/models/hierarchy_vqgan_model.py15index 4b0d657..0bf4712 10064416--- a/models/hierarchy_vqgan_model.py17+++ b/models/hierarchy_vqgan_model.py18@@ -20,7 +20,7 @@ class HierarchyVQSpatialTextureAwareModel():19 20     def __init__(self, opt):21         self.opt = opt22-        self.device = torch.device('cuda')23+        self.device = torch.device(opt['device'])24         self.top_encoder = Encoder(25             ch=opt['top_ch'],26             num_res_blocks=opt['top_num_res_blocks'],27diff --git a/models/parsing_gen_model.py b/models/parsing_gen_model.py28index 9440345..15a1ecb 10064429--- a/models/parsing_gen_model.py30+++ b/models/parsing_gen_model.py31@@ -22,7 +22,7 @@ class ParsingGenModel():32 33     def __init__(self, opt):34         self.opt = opt35-        self.device = torch.device('cuda')36+        self.device = torch.device(opt['device'])37         self.is_train = opt['is_train']38 39         self.attr_embedder = ShapeAttrEmbedding(40diff --git a/models/sample_model.py b/models/sample_model.py41index 4c60e3f..5265cd0 10064442--- a/models/sample_model.py43+++ b/models/sample_model.py44@@ -23,7 +23,7 @@ class BaseSampleModel():45 46     def __init__(self, opt):47         self.opt = opt48-        self.device = torch.device('cuda')49+        self.device = torch.device(opt['device'])50 51         # hierarchical VQVAE52         self.decoder = Decoder(53@@ -123,7 +123,7 @@ class BaseSampleModel():54 55     def load_top_pretrain_models(self):56         # load pretrained vqgan57-        top_vae_checkpoint = torch.load(self.opt['top_vae_path'])58+        top_vae_checkpoint = torch.load(self.opt['top_vae_path'], map_location=self.device)59 60         self.decoder.load_state_dict(61             top_vae_checkpoint['decoder'], strict=True)62@@ -137,7 +137,7 @@ class BaseSampleModel():63         self.top_post_quant_conv.eval()64 65     def load_bot_pretrain_network(self):66-        checkpoint = torch.load(self.opt['bot_vae_path'])67+        checkpoint = torch.load(self.opt['bot_vae_path'], map_location=self.device)68         self.bot_decoder_res.load_state_dict(69             checkpoint['bot_decoder_res'], strict=True)70         self.decoder.load_state_dict(checkpoint['decoder'], strict=True)71@@ -153,7 +153,7 @@ class BaseSampleModel():72 73     def load_pretrained_segm_token(self):74         # load pretrained vqgan for segmentation mask75-        segm_token_checkpoint = torch.load(self.opt['segm_token_path'])76+        segm_token_checkpoint = torch.load(self.opt['segm_token_path'], map_location=self.device)77         self.segm_encoder.load_state_dict(78             segm_token_checkpoint['encoder'], strict=True)79         self.segm_quantizer.load_state_dict(80@@ -166,7 +166,7 @@ class BaseSampleModel():81         self.segm_quant_conv.eval()82 83     def load_index_pred_network(self):84-        checkpoint = torch.load(self.opt['pretrained_index_network'])85+        checkpoint = torch.load(self.opt['pretrained_index_network'], map_location=self.device)86         self.index_pred_guidance_encoder.load_state_dict(87             checkpoint['guidance_encoder'], strict=True)88         self.index_pred_decoder.load_state_dict(89@@ -176,7 +176,7 @@ class BaseSampleModel():90         self.index_pred_decoder.eval()91 92     def load_sampler_pretrained_network(self):93-        checkpoint = torch.load(self.opt['pretrained_sampler'])94+        checkpoint = torch.load(self.opt['pretrained_sampler'], map_location=self.device)95         self.sampler_fn.load_state_dict(checkpoint, strict=True)96         self.sampler_fn.eval()97 98@@ -397,7 +397,7 @@ class SampleFromPoseModel(BaseSampleModel):99                         [185, 210, 205], [130, 165, 180], [225, 141, 151]]100 101     def load_shape_generation_models(self):102-        checkpoint = torch.load(self.opt['pretrained_parsing_gen'])103+        checkpoint = torch.load(self.opt['pretrained_parsing_gen'], map_location=self.device)104 105         self.shape_attr_embedder.load_state_dict(106             checkpoint['embedder'], strict=True)107diff --git a/models/transformer_model.py b/models/transformer_model.py108index 7db0f3e..4523d17 100644109--- a/models/transformer_model.py110+++ b/models/transformer_model.py111@@ -21,7 +21,7 @@ class TransformerTextureAwareModel():112 113     def __init__(self, opt):114         self.opt = opt115-        self.device = torch.device('cuda')116+        self.device = torch.device(opt['device'])117         self.is_train = opt['is_train']118 119         # VQVAE for image120@@ -317,10 +317,10 @@ class TransformerTextureAwareModel():121     def sample_fn(self, temp=1.0, sample_steps=None):122         self._denoise_fn.eval()123 124-        b, device = self.image.size(0), 'cuda'125+        b = self.image.size(0)126         x_t = torch.ones(127-            (b, np.prod(self.shape)), device=device).long() * self.mask_id128-        unmasked = torch.zeros_like(x_t, device=device).bool()129+            (b, np.prod(self.shape)), device=self.device).long() * self.mask_id130+        unmasked = torch.zeros_like(x_t, device=self.device).bool()131         sample_steps = list(range(1, sample_steps + 1))132 133         texture_mask_flatten = self.texture_tokens.view(-1)134@@ -336,11 +336,11 @@ class TransformerTextureAwareModel():135 136         for t in reversed(sample_steps):137             print(f'Sample timestep {t:4d}', end='\r')138-            t = torch.full((b, ), t, device=device, dtype=torch.long)139+            t = torch.full((b, ), t, device=self.device, dtype=torch.long)140 141             # where to unmask142             changes = torch.rand(143-                x_t.shape, device=device) < 1 / t.float().unsqueeze(-1)144+                x_t.shape, device=self.device) < 1 / t.float().unsqueeze(-1)145             # don't unmask somewhere already unmasked146             changes = torch.bitwise_xor(changes,147                                         torch.bitwise_and(changes, unmasked))148diff --git a/models/vqgan_model.py b/models/vqgan_model.py149index 13a2e70..9c840f1 100644150--- a/models/vqgan_model.py151+++ b/models/vqgan_model.py152@@ -20,7 +20,7 @@ class VQModel():153     def __init__(self, opt):154         super().__init__()155         self.opt = opt156-        self.device = torch.device('cuda')157+        self.device = torch.device(opt['device'])158         self.encoder = Encoder(159             ch=opt['ch'],160             num_res_blocks=opt['num_res_blocks'],161@@ -390,7 +390,7 @@ class VQImageSegmTextureModel(VQImageModel):162 163     def __init__(self, opt):164         self.opt = opt165-        self.device = torch.device('cuda')166+        self.device = torch.device(opt['device'])167         self.encoder = Encoder(168             ch=opt['ch'],169             num_res_blocks=opt['num_res_blocks'],170