From c2ab0b11c34acd02d222db14642738ff1f6856cf Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 14 Oct 2024 09:29:33 -0400 Subject: [PATCH] check te device Signed-off-by: Vladimir Mandic --- modules/devices.py | 10 ++++++++++ modules/prompt_parser_diffusers.py | 13 +++++++++++-- modules/sd_hijack_accelerate.py | 14 ++------------ 3 files changed, 23 insertions(+), 14 deletions(-) diff --git a/modules/devices.py b/modules/devices.py index 1269f4eea..f00143d50 100644 --- a/modules/devices.py +++ b/modules/devices.py @@ -505,3 +505,13 @@ def test_for_nans(x, where): message = "A tensor with all NaNs was produced." message += " Use --disable-nan-check commandline argument to disable this check." raise NansException(message) + + +def same_device(d1, d2): + if d1.type != d2.type: + return False + if d1.type == "cuda" and d1.index is None: + d1 = torch.device("cuda", index=0) + if d2.type == "cuda" and d2.index is None: + d2 = torch.device("cuda", index=0) + return d1 == d2 diff --git a/modules/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index c063833a2..599ddf815 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -18,6 +18,8 @@ cache = {} def compel_hijack(self, token_ids: torch.Tensor, attention_mask: typing.Optional[torch.Tensor] = None) -> torch.Tensor: + if not devices.same_device(self.text_encoder.device, devices.device): + sd_models.move_model(self.text_encoder, devices.device) needs_hidden_states = self.returned_embeddings_type != 1 text_encoder_output = self.text_encoder(token_ids, attention_mask, output_hidden_states=needs_hidden_states, return_dict=True) @@ -298,13 +300,20 @@ def prepare_embedding_providers(pipe, clip_skip) -> list[EmbeddingsProvider]: no_mask_provider = EmbeddingsProvider(padding_attention_mask_value=1 if "sote" in pipe.sd_checkpoint_info.name.lower() else 0, tokenizer=pipe.prior_pipe.tokenizer, text_encoder=pipe.prior_pipe.text_encoder, truncate=False, returned_embeddings_type=embedding_type, device=device) embeddings_providers.append(no_mask_provider) elif getattr(pipe, "tokenizer", None) is not None and getattr(pipe, "text_encoder", None) is not None: - sd_models.move_model(pipe.text_encoder, device) + if not devices.same_device(pipe.text_encoder.device, devices.device): + sd_models.move_model(pipe.text_encoder, devices.device) provider = EmbeddingsProvider(tokenizer=pipe.tokenizer, text_encoder=pipe.text_encoder, truncate=False, returned_embeddings_type=embedding_type, device=device) embeddings_providers.append(provider) if getattr(pipe, "tokenizer_2", None) is not None and getattr(pipe, "text_encoder_2", None) is not None: - sd_models.move_model(pipe.text_encoder, device) + if not devices.same_device(pipe.text_encoder_2.device, devices.device): + sd_models.move_model(pipe.text_encoder_2, devices.device) provider = EmbeddingsProvider(tokenizer=pipe.tokenizer_2, text_encoder=pipe.text_encoder_2, truncate=False, returned_embeddings_type=embedding_type, device=device) embeddings_providers.append(provider) + if getattr(pipe, "tokenizer_3", None) is not None and getattr(pipe, "text_encoder_3", None) is not None: + if not devices.same_device(pipe.text_encoder_3.device, devices.device): + sd_models.move_model(pipe.text_encoder_3, devices.device) + provider = EmbeddingsProvider(tokenizer=pipe.tokenizer_3, text_encoder=pipe.text_encoder_3, truncate=False, returned_embeddings_type=embedding_type, device=device) + embeddings_providers.append(provider) return embeddings_providers diff --git a/modules/sd_hijack_accelerate.py b/modules/sd_hijack_accelerate.py index 9eed593ed..90eac5c4e 100644 --- a/modules/sd_hijack_accelerate.py +++ b/modules/sd_hijack_accelerate.py @@ -13,16 +13,6 @@ orig_set_module = accelerate.utils.set_module_tensor_to_device orig_torch_conv = torch.nn.modules.conv.Conv2d._conv_forward # pylint: disable=protected-access -def check_device_same(d1, d2): - if d1.type != d2.type: - return False - if d1.type == "cuda" and d1.index is None: - d1 = torch.device("cuda", index=0) - if d2.type == "cuda" and d2.index is None: - d2 = torch.device("cuda", index=0) - return d1 == d2 - - # called for every item in state_dict by diffusers during model load def hijack_set_module_tensor( module: nn.Module, @@ -46,7 +36,7 @@ def hijack_set_module_tensor( # note: majority of time is spent on .to(old_value.dtype) if tensor_name in module._buffers: # pylint: disable=protected-access module._buffers[tensor_name] = value.to(device, old_value.dtype, non_blocking=True) # pylint: disable=protected-access - elif value is not None or not check_device_same(torch.device(device), module._parameters[tensor_name].device): # pylint: disable=protected-access + elif value is not None or not devices.same_device(torch.device(device), module._parameters[tensor_name].device): # pylint: disable=protected-access param_cls = type(module._parameters[tensor_name]) # pylint: disable=protected-access module._parameters[tensor_name] = param_cls(value, requires_grad=old_value.requires_grad).to(device, old_value.dtype, non_blocking=True) # pylint: disable=protected-access t1 = time.time() @@ -74,7 +64,7 @@ def hijack_set_module_tensor_simple( with devices.inference_context(): if tensor_name in module._buffers: # pylint: disable=protected-access module._buffers[tensor_name] = value.to(device, non_blocking=True) # pylint: disable=protected-access - elif value is not None or not check_device_same(torch.device(device), module._parameters[tensor_name].device): # pylint: disable=protected-access + elif value is not None or not devices.same_device(torch.device(device), module._parameters[tensor_name].device): # pylint: disable=protected-access param_cls = type(module._parameters[tensor_name]) # pylint: disable=protected-access module._parameters[tensor_name] = param_cls(value, requires_grad=old_value.requires_grad).to(device, non_blocking=True) # pylint: disable=protected-access t1 = time.time()