(self)
| 56 | return getattr(self.diffusers_pipeline, name) |
| 57 | |
| 58 | def load_diffusion_model(self): |
| 59 | dtype = self.model_config['dtype'] |
| 60 | transformer_dtype = self.model_config.get('transformer_dtype', dtype) |
| 61 | |
| 62 | with open('configs/auraflow/transformer_config.json') as f: |
| 63 | json_config = json.load(f) |
| 64 | with init_empty_weights(): |
| 65 | transformer = diffusers.AuraFlowTransformer2DModel.from_config(json_config) |
| 66 | state_dict = { |
| 67 | k.replace('model.', ''): v |
| 68 | for k, v in load_state_dict(self.model_config['transformer_path']).items() |
| 69 | } |
| 70 | state_dict = diffusers.loaders.single_file_utils.convert_auraflow_transformer_checkpoint_to_diffusers(state_dict) |
| 71 | for key, tensor in state_dict.items(): |
| 72 | dtype_to_use = dtype if any(keyword in key for keyword in KEEP_IN_HIGH_PRECISION) or tensor.ndim == 1 else transformer_dtype |
| 73 | set_module_tensor_to_device(transformer, key, device='cpu', dtype=dtype_to_use, value=tensor) |
| 74 | |
| 75 | transformer.train() |
| 76 | self.diffusers_pipeline.transformer = transformer |
| 77 | |
| 78 | def get_vae(self): |
| 79 | return self.vae |
nothing calls this directly
no test coverage detected