mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
[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:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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`."
|
||||
|
||||
Reference in New Issue
Block a user