building torch optimizers with simple configs
logger
module-attribute
logger = getLogger('SaigeResearch')
_types
module-attribute
_types = {(__name__): _type for _type in _types}
build_optimizer
build_optimizer(model_params: Iterator[Parameter], _target_: str = 'SGD', **params) -> Optimizer
build optimizer object
Parameters:
-
model_params
(Iterator[Parameter])
–
parameters of model to be trained
-
_target_
(str, default:
'SGD'
)
–
optimizer name. Defaults to "SGD".
Raises:
Returns:
-
Optimizer ( Optimizer
) –
Source code in SaigeToolkit/learning/optimizer.py
| def build_optimizer(model_params: Iterator[Parameter], _target_: str = "SGD", **params) -> Optimizer:
"""build optimizer object
Args:
model_params (Iterator[Parameter]): parameters of model to be trained
_target_ (str, optional): optimizer name. Defaults to "SGD".
Raises:
NotImplementedError: _description_
Returns:
Optimizer: optimizer object
"""
if _target_ not in _types:
raise NotImplementedError(f"OPTIMIZER {_target_} not implemented")
logger.info(f"[{'OPTIMIZER'.center(9)}] {_target_} [params] {params}")
return _types[_target_](model_params, **params)
|