(self)
| 69 | te_wrapper._load_fn = new_load_fn |
| 70 | |
| 71 | def load_diffusion_model(self): |
| 72 | if self.is_32b: |
| 73 | # Model is so big, it's easy to OOM while loading with multiple GPUs. |
| 74 | with one_at_a_time(): |
| 75 | rank = int(os.environ['LOCAL_RANK']) |
| 76 | print(f'Loading Flux2 on rank {rank}') |
| 77 | super().load_diffusion_model() |
| 78 | else: |
| 79 | super().load_diffusion_model() |
| 80 | |
| 81 | def get_call_vae_fn(self, vae): |
| 82 | """ |
nothing calls this directly
no test coverage detected