(self)
| 146 | return getattr(self.diffusers_pipeline, name) |
| 147 | |
| 148 | def load_diffusion_model(self): |
| 149 | dtype = self.model_config['dtype'] |
| 150 | transformer_dtype = self.model_config.get('transformer_dtype', dtype) |
| 151 | with init_empty_weights(): |
| 152 | transformer = Chroma(chroma_params) |
| 153 | transformer.load_state_dict(load_state_dict(self.model_config['transformer_path']), assign=True) |
| 154 | self.diffusers_pipeline.transformer = transformer |
| 155 | |
| 156 | if 'adapter' in self.config: |
| 157 | if fuse_adapters := self.config['adapter'].get('fuse_adapters', None): |
| 158 | print(f'Fusing adapters: {fuse_adapters}') |
| 159 | for fuse_adapter in fuse_adapters: |
| 160 | self.load_and_fuse_adapter(fuse_adapter['path']) |
| 161 | |
| 162 | for name, p in transformer.named_parameters(): |
| 163 | if not any(x in name for x in KEEP_IN_HIGH_PRECISION): |
| 164 | p.data = p.data.to(transformer_dtype) |
| 165 | |
| 166 | self.transformer.train() |
| 167 | for name, p in self.transformer.named_parameters(): |
| 168 | p.original_name = name |
| 169 | |
| 170 | def get_vae(self): |
| 171 | return self.vae |
nothing calls this directly
no test coverage detected