(self)
| 41 | return getattr(self.diffusers_pipeline, name) |
| 42 | |
| 43 | def load_diffusion_model(self): |
| 44 | dtype = self.model_config['dtype'] |
| 45 | transformer_dtype = self.model_config.get('transformer_dtype', dtype) |
| 46 | if diffusers_path := self.model_config.get('diffusers_path', None): |
| 47 | transformer = diffusers.SD3Transformer2DModel.from_pretrained(diffusers_path, torch_dtype=dtype, subfolder='transformer') |
| 48 | else: |
| 49 | raise NotImplementedError() |
| 50 | |
| 51 | for name, p in transformer.named_parameters(): |
| 52 | if not (any(x in name for x in KEEP_IN_HIGH_PRECISION) or p.ndim == 1): |
| 53 | p.data = p.data.to(transformer_dtype) |
| 54 | |
| 55 | self.diffusers_pipeline.transformer = transformer |
| 56 | |
| 57 | def get_vae(self): |
| 58 | return self.vae |
no test coverage detected