(self)
| 56 | self.model_ema_config = self.cfg.get('DIFFUSION_MODEL_EMA', None) |
| 57 | |
| 58 | def construct_network(self): |
| 59 | # embedding_context = torch.device("meta") if self.model_config.get("PRETRAINED_MODEL", None) else nullcontext() |
| 60 | # with embedding_context: |
| 61 | self.model = BACKBONES.build(self.model_config, logger=self.logger).to(torch.bfloat16) |
| 62 | self.logger.info('all parameters:{}'.format(count_params(self.model))) |
| 63 | if self.use_ema: |
| 64 | if self.model_ema_config: |
| 65 | self.model_ema = BACKBONES.build(self.model_ema_config, |
| 66 | logger=self.logger) |
| 67 | else: |
| 68 | self.model_ema = copy.deepcopy(self.model) |
| 69 | self.model_ema = self.model_ema.eval() |
| 70 | for param in self.model_ema.parameters(): |
| 71 | param.requires_grad = False |
| 72 | if self.loss_config: |
| 73 | self.loss = LOSSES.build(self.loss_config, logger=self.logger) |
| 74 | if self.tokenizer_config is not None: |
| 75 | self.tokenizer = TOKENIZERS.build(self.tokenizer_config, |
| 76 | logger=self.logger) |
| 77 | if self.first_stage_config: |
| 78 | self.first_stage_model = MODELS.build(self.first_stage_config, |
| 79 | logger=self.logger) |
| 80 | self.first_stage_model = self.first_stage_model.eval() |
| 81 | self.first_stage_model.train = disabled_train |
| 82 | for param in self.first_stage_model.parameters(): |
| 83 | param.requires_grad = False |
| 84 | else: |
| 85 | self.first_stage_model = None |
| 86 | if self.tokenizer_config is not None: |
| 87 | self.cond_stage_config.KWARGS = { |
| 88 | 'vocab_size': self.tokenizer.vocab_size |
| 89 | } |
| 90 | if self.cond_stage_config == '__is_unconditional__': |
| 91 | print( |
| 92 | f'Training {self.__class__.__name__} as an unconditional model.' |
| 93 | ) |
| 94 | self.cond_stage_model = None |
| 95 | else: |
| 96 | model = EMBEDDERS.build(self.cond_stage_config, logger=self.logger) |
| 97 | self.cond_stage_model = model.eval().requires_grad_(False) |
| 98 | self.cond_stage_model.train = disabled_train |
| 99 | |
| 100 | @torch.no_grad() |
| 101 | def encode_first_stage(self, x, **kwargs): |
nothing calls this directly
no outgoing calls
no test coverage detected