Build an optimizer based on the configuration. Dynamically imports and instantiates an optimizer class from the specified module. Args: parameters: Model parameters to optimize config: FSDPOptimizerConfig with optimizer settings Returns: Optimizer instance
(parameters, config: FSDPOptimizerConfig)
| 144 | |
| 145 | |
| 146 | def build_optimizer(parameters, config: FSDPOptimizerConfig): |
| 147 | """Build an optimizer based on the configuration. |
| 148 | |
| 149 | Dynamically imports and instantiates an optimizer class from the specified module. |
| 150 | |
| 151 | Args: |
| 152 | parameters: Model parameters to optimize |
| 153 | config: FSDPOptimizerConfig with optimizer settings |
| 154 | |
| 155 | Returns: |
| 156 | Optimizer instance |
| 157 | |
| 158 | Examples: |
| 159 | # PyTorch AdamW |
| 160 | config.optimizer_impl = "torch.optim" |
| 161 | config.optimizer = "AdamW" |
| 162 | |
| 163 | # TorchAO AdamW with bf16 stochastic rounding |
| 164 | config.optimizer_impl = "torchao.optim" |
| 165 | config.optimizer = "_AdamW" |
| 166 | config.override_optimizer_config = {"bf16_stochastic_round": True} |
| 167 | |
| 168 | # BitsAndBytes AdamW 8bit |
| 169 | config.optimizer_impl = "bitsandbytes.optim" |
| 170 | config.optimizer = "AdamW8bit" |
| 171 | """ |
| 172 | import importlib |
| 173 | |
| 174 | optimizer_args = { |
| 175 | "lr": config.lr, |
| 176 | "weight_decay": config.weight_decay, |
| 177 | } |
| 178 | |
| 179 | optimizer_name_lower = config.optimizer.lower() |
| 180 | if "adam" in optimizer_name_lower or "ademamix" in optimizer_name_lower: |
| 181 | optimizer_args["betas"] = config.betas |
| 182 | |
| 183 | if config.override_optimizer_config is not None: |
| 184 | optimizer_args.update(config.override_optimizer_config) |
| 185 | |
| 186 | try: |
| 187 | module = importlib.import_module(config.optimizer_impl) |
| 188 | optimizer_cls = getattr(module, config.optimizer) |
| 189 | except ImportError as e: |
| 190 | raise ImportError( |
| 191 | f"Failed to import module '{config.optimizer_impl}'. Make sure the package is installed. Error: {e}" |
| 192 | ) from e |
| 193 | except AttributeError as e: |
| 194 | raise AttributeError( |
| 195 | f"Optimizer '{config.optimizer}' not found in module '{config.optimizer_impl}'. " |
| 196 | f"Available optimizers: {dir(module)}" |
| 197 | ) from e |
| 198 | |
| 199 | return optimizer_cls(parameters, **optimizer_args) |
no test coverage detected