diff --git a/README.md b/README.md index 2dd6e67..59bac73 100644 --- a/README.md +++ b/README.md @@ -394,6 +394,7 @@ If you would like to use a custom dataset, refer to the dataset preparation guid > - Use a DeepSpeed config to launch training (refer to [`accelerate_configs/deepspeed.yaml`](./accelerate_configs/deepspeed.yaml) as an example). > - Pass `--precompute_conditions` when launching training. > - Pass `--gradient_checkpointing` when launching training. +> - Pass `--use_8bit_bnb` when launching training. Note that this is only applicable to Adam and AdamW optimizers. > - Do not perform validation/testing. This saves a significant amount of memory, which can be used to focus solely on training if you're on smaller VRAM GPUs. ## Memory requirements diff --git a/finetrainers/args.py b/finetrainers/args.py index b32a7a5..247d11a 100644 --- a/finetrainers/args.py +++ b/finetrainers/args.py @@ -79,6 +79,7 @@ class Args: # Optimizer arguments optimizer: str = "adamw" + use_8bit_bnb: bool = False lr: float = 1e-4 scale_lr: bool = False lr_scheduler: str = "cosine_with_restarts" @@ -166,6 +167,7 @@ class Args: }, "optimizer_arguments": { "optimizer": self.optimizer, + "use_8bit_bnb": self.use_8bit_bnb, "lr": self.lr, "scale_lr": self.scale_lr, "lr_scheduler": self.lr_scheduler, @@ -532,6 +534,11 @@ def _add_optimizer_arguments(parser: argparse.ArgumentParser) -> None: choices=["adam", "adamw"], help=("The optimizer type to use."), ) + parser.add_argument( + "--use_8bit_bnb", + action="store_true", + help=("Whether to use 8bit variant of the `--optimizer` using `bitsandbytes`."), + ) parser.add_argument( "--beta1", type=float, @@ -741,6 +748,7 @@ def _map_to_args_type(args: Dict[str, Any]) -> Args: # Optimizer arguments result_args.optimizer = args.optimizer or "adamw" + result_args.use_8bit_bnb = args.use_8bit_bnb result_args.lr = args.lr or 1e-4 result_args.scale_lr = args.scale_lr result_args.lr_scheduler = args.lr_scheduler diff --git a/finetrainers/trainer.py b/finetrainers/trainer.py index bbacc49..5523999 100644 --- a/finetrainers/trainer.py +++ b/finetrainers/trainer.py @@ -506,6 +506,7 @@ class Trainer: beta3=self.args.beta3, epsilon=self.args.epsilon, weight_decay=self.args.weight_decay, + use_8bit=self.args.use_8bit_bnb, use_deepspeed=use_deepspeed_opt, ) @@ -1003,7 +1004,7 @@ class Trainer: if self.args.push_to_hub: repo_id = self.args.hub_model_id or Path(self.args.output_dir).name - self.state.repo_id = create_repo(token=self.args.hub_token, repo_id=repo_id).repo_id + self.state.repo_id = create_repo(token=self.args.hub_token, repo_id=repo_id, exist_ok=True).repo_id def _move_components_to_device(self): if self.text_encoder is not None: diff --git a/finetrainers/utils/optimizer_utils.py b/finetrainers/utils/optimizer_utils.py index 77842eb..e988687 100644 --- a/finetrainers/utils/optimizer_utils.py +++ b/finetrainers/utils/optimizer_utils.py @@ -39,6 +39,7 @@ def get_optimizer( weight_decay=weight_decay, ) + # TODO: consider moving the validation logic to `args.py` when we have torchao. if use_8bit and use_4bit: raise ValueError("Cannot set both `use_8bit` and `use_4bit` to True.") @@ -46,7 +47,6 @@ def get_optimizer( try: import torchao - torchao.__version__ except ImportError: raise ImportError( "To use optimizers from torchao, please install the torchao library: `USE_CPP=0 pip install torchao`."