radames/Text2Human-API
1
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 