2022-09-08 03:59:30 +00:00
|
|
|
import importlib
|
2022-09-09 04:51:25 +00:00
|
|
|
import logging
|
2022-09-11 06:27:22 +00:00
|
|
|
import platform
|
2022-09-22 05:03:12 +00:00
|
|
|
from contextlib import contextmanager, nullcontext
|
2022-09-08 03:59:30 +00:00
|
|
|
from functools import lru_cache
|
2022-10-04 22:07:40 +00:00
|
|
|
from typing import Any, List, Optional, Union
|
2022-09-08 03:59:30 +00:00
|
|
|
|
|
|
|
import torch
|
2022-09-22 05:03:12 +00:00
|
|
|
from torch import Tensor, autocast
|
2022-09-16 16:24:24 +00:00
|
|
|
from torch.nn import functional
|
2022-09-13 07:46:37 +00:00
|
|
|
from torch.overrides import handle_torch_function, has_torch_function_variadic
|
2022-09-08 03:59:30 +00:00
|
|
|
|
2022-09-09 04:51:25 +00:00
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
2022-09-08 03:59:30 +00:00
|
|
|
|
|
|
|
@lru_cache()
|
2022-10-04 22:07:40 +00:00
|
|
|
def get_device() -> str:
|
|
|
|
"""Return the best torch backend available"""
|
2022-09-08 03:59:30 +00:00
|
|
|
if torch.cuda.is_available():
|
|
|
|
return "cuda"
|
2022-09-16 16:24:24 +00:00
|
|
|
|
|
|
|
if torch.backends.mps.is_available():
|
2022-09-17 19:24:27 +00:00
|
|
|
return "mps:0"
|
2022-09-16 16:24:24 +00:00
|
|
|
|
|
|
|
return "cpu"
|
2022-09-08 03:59:30 +00:00
|
|
|
|
|
|
|
|
2022-09-11 06:27:22 +00:00
|
|
|
@lru_cache()
|
2022-10-04 22:07:40 +00:00
|
|
|
def get_hardware_description(device_type: str) -> str:
|
|
|
|
"""Description of the hardware being used"""
|
|
|
|
desc = platform.platform()
|
2022-09-11 06:27:22 +00:00
|
|
|
if device_type == "cuda":
|
2022-10-04 22:07:40 +00:00
|
|
|
desc += "-" + torch.cuda.get_device_name(0)
|
2022-09-11 06:27:22 +00:00
|
|
|
|
2022-10-04 22:07:40 +00:00
|
|
|
return desc
|
2022-09-11 06:27:22 +00:00
|
|
|
|
2022-09-08 03:59:30 +00:00
|
|
|
|
2022-10-04 22:07:40 +00:00
|
|
|
def get_obj_from_str(import_path: str, reload=False) -> Any:
|
|
|
|
"""
|
|
|
|
Gets a python object from a string reference if it's location
|
|
|
|
|
|
|
|
Example: "functools.lru_cache"
|
|
|
|
"""
|
|
|
|
module_path, obj_name = import_path.rsplit(".", 1)
|
|
|
|
if reload:
|
|
|
|
module_imp = importlib.import_module(module_path)
|
|
|
|
importlib.reload(module_imp)
|
|
|
|
module = importlib.import_module(module_path, package=None)
|
|
|
|
return getattr(module, obj_name)
|
2022-09-08 03:59:30 +00:00
|
|
|
|
2022-10-04 22:07:40 +00:00
|
|
|
|
|
|
|
def instantiate_from_config(config: Union[dict, str]) -> Any:
|
|
|
|
"""Instantiate an object from a config dict"""
|
2022-09-16 16:24:24 +00:00
|
|
|
if "target" not in config:
|
2022-09-08 03:59:30 +00:00
|
|
|
if config == "__is_first_stage__":
|
|
|
|
return None
|
2022-09-16 16:24:24 +00:00
|
|
|
if config == "__is_unconditional__":
|
2022-09-08 03:59:30 +00:00
|
|
|
return None
|
|
|
|
raise KeyError("Expected key `target` to instantiate.")
|
2022-10-04 22:07:40 +00:00
|
|
|
params = config.get("params", {})
|
|
|
|
_cls = get_obj_from_str(config["target"])
|
|
|
|
return _cls(**params)
|
2022-09-10 07:32:31 +00:00
|
|
|
|
|
|
|
|
2022-09-22 05:38:44 +00:00
|
|
|
@contextmanager
|
|
|
|
def platform_appropriate_autocast(precision="autocast"):
|
|
|
|
"""
|
2022-09-24 05:58:48 +00:00
|
|
|
Allow calculations to run in mixed precision, which can be faster
|
2022-09-22 05:38:44 +00:00
|
|
|
"""
|
|
|
|
precision_scope = nullcontext
|
|
|
|
if precision == "autocast" and get_device() in ("cuda", "cpu"):
|
|
|
|
precision_scope = autocast
|
|
|
|
with precision_scope(get_device()):
|
|
|
|
yield
|
|
|
|
|
|
|
|
|
2022-09-10 07:32:31 +00:00
|
|
|
def _fixed_layer_norm(
|
2022-09-16 16:24:24 +00:00
|
|
|
input: Tensor, # noqa
|
2022-09-10 07:32:31 +00:00
|
|
|
normalized_shape: List[int],
|
|
|
|
weight: Optional[Tensor] = None,
|
|
|
|
bias: Optional[Tensor] = None,
|
|
|
|
eps: float = 1e-5,
|
|
|
|
) -> Tensor:
|
2022-09-16 16:24:24 +00:00
|
|
|
"""
|
|
|
|
Applies Layer Normalization for last certain number of dimensions.
|
|
|
|
|
2022-09-10 07:32:31 +00:00
|
|
|
See :class:`~torch.nn.LayerNorm` for details.
|
|
|
|
"""
|
|
|
|
if has_torch_function_variadic(input, weight, bias):
|
|
|
|
return handle_torch_function(
|
|
|
|
_fixed_layer_norm,
|
|
|
|
(input, weight, bias),
|
|
|
|
input,
|
|
|
|
normalized_shape,
|
|
|
|
weight=weight,
|
|
|
|
bias=bias,
|
|
|
|
eps=eps,
|
|
|
|
)
|
|
|
|
return torch.layer_norm(
|
|
|
|
input.contiguous(),
|
|
|
|
normalized_shape,
|
|
|
|
weight,
|
|
|
|
bias,
|
|
|
|
eps,
|
|
|
|
torch.backends.cudnn.enabled,
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
@contextmanager
|
|
|
|
def fix_torch_nn_layer_norm():
|
|
|
|
"""https://github.com/CompVis/stable-diffusion/issues/25#issuecomment-1221416526"""
|
|
|
|
orig_function = functional.layer_norm
|
|
|
|
functional.layer_norm = _fixed_layer_norm
|
|
|
|
try:
|
|
|
|
yield
|
|
|
|
finally:
|
|
|
|
functional.layer_norm = orig_function
|
2022-09-12 01:00:40 +00:00
|
|
|
|
|
|
|
|
2022-09-22 05:03:12 +00:00
|
|
|
@contextmanager
|
|
|
|
def fix_torch_group_norm():
|
|
|
|
"""
|
|
|
|
Patch group_norm to cast the weights to the same type as the inputs
|
|
|
|
|
|
|
|
From what I can understand all the other repos just switch to full precision instead
|
|
|
|
of addressing this. I think this would make things slower but I'm not sure.
|
|
|
|
|
|
|
|
https://github.com/pytorch/pytorch/pull/81852
|
|
|
|
|
|
|
|
"""
|
|
|
|
|
|
|
|
orig_group_norm = functional.group_norm
|
|
|
|
|
|
|
|
def _group_norm_wrapper(
|
2022-09-22 05:38:44 +00:00
|
|
|
input: Tensor, # noqa
|
2022-09-22 05:03:12 +00:00
|
|
|
num_groups: int,
|
|
|
|
weight: Optional[Tensor] = None,
|
|
|
|
bias: Optional[Tensor] = None,
|
|
|
|
eps: float = 1e-5,
|
|
|
|
) -> Tensor:
|
|
|
|
if weight is not None and weight.dtype != input.dtype:
|
|
|
|
weight = weight.to(input.dtype)
|
|
|
|
if bias is not None and bias.dtype != input.dtype:
|
|
|
|
bias = bias.to(input.dtype)
|
|
|
|
|
|
|
|
return orig_group_norm(
|
|
|
|
input=input, num_groups=num_groups, weight=weight, bias=bias, eps=eps
|
|
|
|
)
|
|
|
|
|
|
|
|
functional.group_norm = _group_norm_wrapper
|
|
|
|
try:
|
|
|
|
yield
|
|
|
|
finally:
|
|
|
|
functional.group_norm = orig_group_norm
|
|
|
|
|
|
|
|
|
2022-10-16 23:42:46 +00:00
|
|
|
def randn_seeded(seed: int, size: List[int]) -> Tensor:
|
|
|
|
"""Generate a random tensor with a given seed"""
|
|
|
|
g_cpu = torch.Generator()
|
|
|
|
g_cpu.manual_seed(seed)
|
|
|
|
noise = torch.randn(
|
|
|
|
size,
|
|
|
|
device="cpu",
|
|
|
|
generator=g_cpu,
|
|
|
|
)
|
|
|
|
return noise
|
|
|
|
|
|
|
|
|
|
|
|
def check_torch_working():
|
|
|
|
"""Check that torch is working"""
|
|
|
|
try:
|
|
|
|
torch.randn(1, device=get_device())
|
|
|
|
except RuntimeError as e:
|
|
|
|
if "CUDA" in str(e):
|
|
|
|
raise RuntimeError(
|
|
|
|
"CUDA is not working. Make sure you have a GPU and CUDA installed."
|
|
|
|
) from e
|
|
|
|
raise e
|