mac_specific.py 5.4 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586
  1. import logging
  2. import torch
  3. import platform
  4. from modules.sd_hijack_utils import CondFunc
  5. from packaging import version
  6. log = logging.getLogger(__name__)
  7. # before torch version 1.13, has_mps is only available in nightly pytorch and macOS 12.3+,
  8. # use check `getattr` and try it for compatibility.
  9. # in torch version 1.13, backends.mps.is_available() and backends.mps.is_built() are introduced in to check mps availabilty,
  10. # since torch 2.0.1+ nightly build, getattr(torch, 'has_mps', False) was deprecated, see https://github.com/pytorch/pytorch/pull/103279
  11. def check_for_mps() -> bool:
  12. if version.parse(torch.__version__) <= version.parse("2.0.1"):
  13. if not getattr(torch, 'has_mps', False):
  14. return False
  15. try:
  16. torch.zeros(1).to(torch.device("mps"))
  17. return True
  18. except Exception:
  19. return False
  20. else:
  21. return torch.backends.mps.is_available() and torch.backends.mps.is_built()
  22. has_mps = check_for_mps()
  23. def torch_mps_gc() -> None:
  24. try:
  25. from modules.shared import state
  26. if state.current_latent is not None:
  27. log.debug("`current_latent` is set, skipping MPS garbage collection")
  28. return
  29. from torch.mps import empty_cache
  30. empty_cache()
  31. except Exception:
  32. log.warning("MPS garbage collection failed", exc_info=True)
  33. # MPS workaround for https://github.com/pytorch/pytorch/issues/89784
  34. def cumsum_fix(input, cumsum_func, *args, **kwargs):
  35. if input.device.type == 'mps':
  36. output_dtype = kwargs.get('dtype', input.dtype)
  37. if output_dtype == torch.int64:
  38. return cumsum_func(input.cpu(), *args, **kwargs).to(input.device)
  39. elif output_dtype == torch.bool or cumsum_needs_int_fix and (output_dtype == torch.int8 or output_dtype == torch.int16):
  40. return cumsum_func(input.to(torch.int32), *args, **kwargs).to(torch.int64)
  41. return cumsum_func(input, *args, **kwargs)
  42. if has_mps:
  43. # MPS fix for randn in torchsde
  44. CondFunc('torchsde._brownian.brownian_interval._randn', lambda _, size, dtype, device, seed: torch.randn(size, dtype=dtype, device=torch.device("cpu"), generator=torch.Generator(torch.device("cpu")).manual_seed(int(seed))).to(device), lambda _, size, dtype, device, seed: device.type == 'mps')
  45. if platform.mac_ver()[0].startswith("13.2."):
  46. # MPS workaround for https://github.com/pytorch/pytorch/issues/95188, thanks to danieldk (https://github.com/explosion/curated-transformers/pull/124)
  47. CondFunc('torch.nn.functional.linear', lambda _, input, weight, bias: (torch.matmul(input, weight.t()) + bias) if bias is not None else torch.matmul(input, weight.t()), lambda _, input, weight, bias: input.numel() > 10485760)
  48. if version.parse(torch.__version__) < version.parse("1.13"):
  49. # PyTorch 1.13 doesn't need these fixes but unfortunately is slower and has regressions that prevent training from working
  50. # MPS workaround for https://github.com/pytorch/pytorch/issues/79383
  51. CondFunc('torch.Tensor.to', lambda orig_func, self, *args, **kwargs: orig_func(self.contiguous(), *args, **kwargs),
  52. lambda _, self, *args, **kwargs: self.device.type != 'mps' and (args and isinstance(args[0], torch.device) and args[0].type == 'mps' or isinstance(kwargs.get('device'), torch.device) and kwargs['device'].type == 'mps'))
  53. # MPS workaround for https://github.com/pytorch/pytorch/issues/80800
  54. CondFunc('torch.nn.functional.layer_norm', lambda orig_func, *args, **kwargs: orig_func(*([args[0].contiguous()] + list(args[1:])), **kwargs),
  55. lambda _, *args, **kwargs: args and isinstance(args[0], torch.Tensor) and args[0].device.type == 'mps')
  56. # MPS workaround for https://github.com/pytorch/pytorch/issues/90532
  57. CondFunc('torch.Tensor.numpy', lambda orig_func, self, *args, **kwargs: orig_func(self.detach(), *args, **kwargs), lambda _, self, *args, **kwargs: self.requires_grad)
  58. elif version.parse(torch.__version__) > version.parse("1.13.1"):
  59. cumsum_needs_int_fix = not torch.Tensor([1,2]).to(torch.device("mps")).equal(torch.ShortTensor([1,1]).to(torch.device("mps")).cumsum(0))
  60. cumsum_fix_func = lambda orig_func, input, *args, **kwargs: cumsum_fix(input, orig_func, *args, **kwargs)
  61. CondFunc('torch.cumsum', cumsum_fix_func, None)
  62. CondFunc('torch.Tensor.cumsum', cumsum_fix_func, None)
  63. CondFunc('torch.narrow', lambda orig_func, *args, **kwargs: orig_func(*args, **kwargs).clone(), None)
  64. # MPS workaround for https://github.com/pytorch/pytorch/issues/96113
  65. CondFunc('torch.nn.functional.layer_norm', lambda orig_func, x, normalized_shape, weight, bias, eps, **kwargs: orig_func(x.float(), normalized_shape, weight.float() if weight is not None else None, bias.float() if bias is not None else bias, eps).to(x.dtype), lambda _, input, *args, **kwargs: len(args) == 4 and input.device.type == 'mps')
  66. # MPS workaround for https://github.com/pytorch/pytorch/issues/92311
  67. if platform.processor() == 'i386':
  68. for funcName in ['torch.argmax', 'torch.Tensor.argmax']:
  69. CondFunc(funcName, lambda _, input, *args, **kwargs: torch.max(input.float() if input.dtype == torch.int64 else input, *args, **kwargs)[1], lambda _, input, *args, **kwargs: input.device.type == 'mps')