INNER CODE UNIT · Python
optimizer_factory
nv-tlabs/ATISS · scene_synthesis/networks/__init__.py:41
def optimizer_factory(config, parameters):
"""Based on the provided config create the suitable optimizer."""
optimizer = config.get("optimizer", "Adam")
lr = config.get("lr", 1e-3)
momentum = config.get("momentum", 0.9)
# weight_decay = config.get("weight_decay", 0.0)
# Weight decay was set to 0.0 in the paper's experiments. We note that
# increasing the weight_decay deteriorates performance.
weight_decay = 0.0
if optimizer == "SGD":
return torch.optim.SGD(
parameters, lr=lr, momentum=momentum, weight_decay=weight_decay
)
elif optimizer == "Adam":
return torch.optim.Adam(parameters, lr=lr, weight_decay=weight_decay)
elif optimizer == "RAdam":
return RAdam(parameters, lr=lr, weight_decay=weight_decay)