diff --git a/qmp/utility/context.py b/qmp/utility/context.py index 5ab948e..a90e00b 100644 --- a/qmp/utility/context.py +++ b/qmp/utility/context.py @@ -199,7 +199,7 @@ def create_optimizer( logging.info("Initializing the optimizer") optimizer_t = getattr(torch.optim, optimizer_config.name) - optimizer = optimizer_t(params=params, **optimizer_config.params) # type: ignore[arg-type] + optimizer: torch.optim.Optimizer = optimizer_t(params=params, **optimizer_config.params) if state_dict is not None: logging.info("Loading state dict of the optimizer") diff --git a/qmp/utility/model_dict.py b/qmp/utility/model_dict.py index d9ddc86..883aa4d 100644 --- a/qmp/utility/model_dict.py +++ b/qmp/utility/model_dict.py @@ -64,18 +64,6 @@ def to(self, device: torch.device | None = None, dtype: torch.dtype | None = Non def parameters(self) -> typing.Iterable[torch.Tensor]: """torch.nn.Module function""" - def bfloat16(self) -> typing.Self: - """torch.nn.Module function""" - - def half(self) -> typing.Self: - """torch.nn.Module function""" - - def float(self) -> typing.Self: - """torch.nn.Module function""" - - def double(self) -> typing.Self: - """torch.nn.Module function""" - Model_contra = typing.TypeVar("Model_contra", contravariant=True)