diff --git a/modules/sd_models.py b/modules/sd_models.py index 94824218a..5091a784b 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -411,7 +411,7 @@ def load_diffuser_force(model_type, checkpoint_info, diffusers_load_config, op=' shared.log.error(f'Load {op}: path="{checkpoint_info.path}" {e}') if debug_load: errors.display(e, 'Load') - return None, True + return None return sd_model diff --git a/modules/sd_te_remote.py b/modules/sd_te_remote.py new file mode 100644 index 000000000..264eeb145 --- /dev/null +++ b/modules/sd_te_remote.py @@ -0,0 +1,40 @@ +from typing import List, Optional, Union +import os +import time +import json +import torch +import requests +from modules import devices, errors + + +def get_t5_prompt_embeds( + prompt: Union[str, List[str]] = None, + num_images_per_prompt: int = 1, # pylint: disable=unused-argument + max_sequence_length: int = 512, # pylint: disable=unused-argument + device: Optional[torch.device] = None, + dtype: Optional[torch.dtype] = None, +): + device = device or devices.device + dtype = dtype or devices.dtype + url = os.environ.get('SD_REMOTE_T5', None) + if url is None: + errors.log.error('Remote-TE: url is not set') + return None + try: + t0 = time.time() + response = requests.post( + url=url, + headers={ "Content-Type": "application/json" }, + json=prompt, + timeout=300, + ) + t1 = time.time() + shape = json.loads(response.headers["shape"]) + buffer = bytearray(response.content) + tensor = torch.frombuffer(buffer, dtype=dtype).reshape(shape) + errors.log.debug(f'Remote-TE: url="{url}" prompt="{prompt}" shape={shape} time={t1-t0:.3f}') + return tensor.to(device=device, dtype=dtype) + except Exception as e: + errors.log.error(f'Remote-TE: {e}') + errors.display(e, 'remote-te') + return None diff --git a/pipelines/model_flux.py b/pipelines/model_flux.py index 34de4da83..3fa0fadd3 100644 --- a/pipelines/model_flux.py +++ b/pipelines/model_flux.py @@ -1,3 +1,4 @@ +import os import diffusers import transformers from modules import shared, devices, sd_models, model_quant @@ -64,6 +65,12 @@ def load_flux(checkpoint_info, diffusers_load_config={}): **load_args, ) + if os.environ.get('SD_REMOTE_T5', None) is not None: + from modules import sd_te_remote + shared.log.warning('Remote-TE: applying patch') + pipe._get_t5_prompt_embeds = sd_te_remote.get_t5_prompt_embeds # pylint: disable=protected-access + pipe.text_encoder_2 = None + del text_encoder_2 del transformer