Skip to content

utils

shrinkai.utils

Functions:

Name Description
resolve_device

Resolves the computation device automatically.

Functions:

resolve_device

resolve_device(
    device: str | device = "auto",
) -> torch.device

Resolves the computation device automatically.

Parameters:

Name Type Description Default
device str | device

Computing target ('auto', 'mps', 'cuda', 'cpu' or torch.device). Defaults to 'auto'.

'auto'

Returns: torch.device: torch device.

Source code in src/shrinkai/utils.py
def resolve_device(device: str | torch.device = "auto") -> torch.device:
    """Resolves the computation device automatically.

    Args:
        device: Computing target ('auto', 'mps', 'cuda', 'cpu' or torch.device).
            Defaults to 'auto'.
    Returns:
        torch.device: torch device."""
    if isinstance(device, str):
        if device == "auto":
            return torch.device(
                "cuda"
                if torch.cuda.is_available()
                else "mps"
                if torch.backends.mps.is_available()
                else "cpu"
            )
        else:
            return torch.device(device)
    else:
        return device