copy conds to device before repeating

This commit is contained in:
drhead 2025-05-25 12:49:49 -04:00 committed by GitHub
parent 87f4673af9
commit baf1e9f1fe
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

View File

@ -11,17 +11,13 @@ class CONDRegular:
def _copy_with(self, cond):
return self.__class__(cond)
def _pin_memory(self, cond):
if cond.device == torch.device('cpu'):
return cond.pin_memory()
else:
return cond
def _pin_cond(self, device):
if self.cond.device == torch.device('cpu') and device_supports_non_blocking(device):
self.cond = self.cond.pin_memory(device)
def process_cond(self, batch_size, device, **kwargs):
if device_supports_non_blocking(device):
return self._copy_with(comfy.utils.repeat_to_batch_size(self._pin_memory(self.cond), batch_size).to(device, non_blocking=True))
else:
return self._copy_with(comfy.utils.repeat_to_batch_size(self.cond, batch_size).to(device))
self._pin_cond(device)
return self._copy_with(comfy.utils.repeat_to_batch_size(self.cond.to(device, non_blocking=device_supports_non_blocking(device)), batch_size))
def can_concat(self, other):
if self.cond.shape != other.cond.shape: