[optimization] support 8bit optims from bistandbytes (#163)

* support 8bit optims from bnb.

* fix exist_ok

* use_8bit_bnb.

* fix.

* note in readme.
This commit is contained in:
Sayak Paul
2024-12-30 14:03:56 +05:30
committed by GitHub
parent e9fba4ab7f
commit 191ea8b0cc
4 changed files with 12 additions and 2 deletions
+1
View File
@@ -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
+8
View File
@@ -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
+2 -1
View File
@@ -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:
+1 -1
View File
@@ -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`."