Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 9 additions & 1 deletion src/diffusers/pipelines/pipeline_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -601,12 +601,20 @@ def module_is_offloaded(module):
def device(self) -> torch.device:
r"""
Returns:
`torch.device`: The torch device on which the pipeline is located.
`torch.device`: The torch device on which the pipeline is located. When components are split across devices
(for example, text encoders on CPU while the denoising backbone runs on an accelerator), the accelerator
device is returned.
"""
module_names, _ = self._get_signature_keys(self)
modules = [getattr(self, n, None) for n in module_names]
modules = [m for m in modules if isinstance(m, torch.nn.Module)]

# Prefer a non-CPU, non-meta component so a split pipeline reports the accelerator it computes on,
# rather than whichever component happens to sort first.
for module in modules:
if module.device.type not in ("cpu", "meta"):
return module.device

for module in modules:
return module.device

Expand Down
28 changes: 28 additions & 0 deletions tests/pipelines/test_pipelines.py
Original file line number Diff line number Diff line change
Expand Up @@ -1944,6 +1944,34 @@ def test_pipe_to(self):
assert sd1.device.type == device_type
assert sd2.device.type == device_type

@require_torch_accelerator
def test_pipe_device_split_across_devices(self):
unet = self.dummy_cond_unet()
scheduler = PNDMScheduler(skip_prk_steps=True)
vae = self.dummy_vae
bert = self.dummy_text_encoder
tokenizer = CLIPTokenizer.from_pretrained("hf-internal-testing/tiny-random-clip")

sd = StableDiffusionPipeline(
unet=unet,
scheduler=scheduler,
vae=vae,
text_encoder=bert,
tokenizer=tokenizer,
safety_checker=None,
feature_extractor=self.dummy_extractor,
)

device_type = torch.device(torch_device).type

# Text encoder stays on CPU while the denoising backbone runs on the accelerator. `text_encoder` sorts
# before `unet`/`vae`, so a first-component rule would report `cpu` here.
sd.unet.to(torch_device)
sd.vae.to(torch_device)

assert sd.text_encoder.device.type == "cpu"
assert sd.device.type == device_type

def test_pipe_same_device_id_offload(self):
unet = self.dummy_cond_unet()
scheduler = PNDMScheduler(skip_prk_steps=True)
Expand Down
Loading