-
Notifications
You must be signed in to change notification settings - Fork 7.4k
Consolidate torch device backend dispatch #14792
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
ef326f3
8e1c6da
1a0c810
394e39f
99d7c38
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -38,7 +38,7 @@ | |
| scale_lora_layers, | ||
| unscale_lora_layers, | ||
| ) | ||
| from ...utils.torch_utils import get_device, is_torch_version, randn_tensor | ||
| from ...utils.torch_utils import randn_tensor | ||
| from ..pipeline_utils import DiffusionPipeline | ||
| from ..pixart_alpha.pipeline_pixart_alpha import ( | ||
| ASPECT_RATIO_512_BIN, | ||
|
|
@@ -1053,15 +1053,9 @@ def __call__( | |
| image = latents | ||
| else: | ||
| latents = latents.to(self.vae.dtype) | ||
| torch_accelerator_module = getattr(torch, get_device(), torch.cuda) | ||
| oom_error = ( | ||
| torch.OutOfMemoryError | ||
| if is_torch_version(">=", "2.5.0") | ||
| else torch_accelerator_module.OutOfMemoryError | ||
| ) | ||
|
Comment on lines
-1056
to
-1061
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Safe because we pin on >=2.6. |
||
| try: | ||
| image = self.vae.decode(latents / self.vae.config.scaling_factor, return_dict=False)[0] | ||
| except oom_error as e: | ||
| except torch.OutOfMemoryError as e: | ||
| warnings.warn( | ||
| f"{e}. \n" | ||
| f"Try to use VAE tiling for large images. For example: \n" | ||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -18,6 +18,7 @@ | |
| is_torch_available, | ||
| logging, | ||
| ) | ||
| from ...utils.torch_utils import get_device | ||
|
|
||
|
|
||
| if is_torch_available() and is_gguf_available(): | ||
|
|
@@ -177,12 +178,7 @@ def _dequantize(self, model): | |
| logger.info( | ||
| "Model was found to be on CPU (could happen as a result of `enable_model_cpu_offload()`). So, moving it to accelerator. After dequantization, will move the model back to CPU again to preserve the previous device." | ||
| ) | ||
| device = ( | ||
| torch.accelerator.current_accelerator() | ||
| if hasattr(torch, "accelerator") | ||
| else torch.cuda.current_device() | ||
| ) | ||
| model.to(device) | ||
| model.to(get_device()) | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Are all the sites of
Collaborator
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Not everywhere. Just the places where we were inlining the device fallback. Can do a follow up pass to find the other places where this is applicable. |
||
|
|
||
| model = _dequantize_gguf_and_restore_linear(model, self.modules_to_not_convert) | ||
| if is_model_on_cpu: | ||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Deadcode since min supported torch version is 2.6