From 2033993f57df3369c3c0947b989ce3588506b626 Mon Sep 17 00:00:00 2001 From: hcsolakoglu <155680432+hcsolakoglu@users.noreply.github.com> Date: Tue, 5 Nov 2024 15:11:37 +0300 Subject: [PATCH 1/2] Add --bnb_optimizer argument to CLI and pass it to Trainer initialization Add `--bnb_optimizer` argument to CLI and pass it to Trainer initialization. * Add `--bnb_optimizer` argument to `parse_args()` function in `src/f5_tts/train/finetune_cli.py`. * Pass `bnb_optimizer` argument to `Trainer` initialization in the `main()` function of `src/f5_tts/train/finetune_cli.py`. --- src/f5_tts/train/finetune_cli.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/src/f5_tts/train/finetune_cli.py b/src/f5_tts/train/finetune_cli.py index 43d766c..13d861c 100644 --- a/src/f5_tts/train/finetune_cli.py +++ b/src/f5_tts/train/finetune_cli.py @@ -55,7 +55,6 @@ def parse_args(): default=None, help="Path to custom tokenizer vocab file (only used if tokenizer = 'custom')", ) - parser.add_argument( "--log_samples", type=bool, @@ -63,6 +62,12 @@ def parse_args(): help="Log inferenced samples per ckpt save steps", ) parser.add_argument("--logger", type=str, default=None, choices=["wandb", "tensorboard"], help="logger") + parser.add_argument( + "--bnb_optimizer", + type=bool, + default=False, + help="Use 8-bit Adam optimizer from bitsandbytes" + ) return parser.parse_args() @@ -147,6 +152,7 @@ def main(): wandb_resume_id=wandb_resume_id, log_samples=args.log_samples, last_per_steps=args.last_per_steps, + bnb_optimizer=args.bnb_optimizer, ) train_dataset = load_dataset(args.dataset_name, tokenizer, mel_spec_kwargs=mel_spec_kwargs) From dbe35da754fe13fb1280f14bb986e6c5237d92b8 Mon Sep 17 00:00:00 2001 From: Yushen CHEN <45333109+SWivid@users.noreply.github.com> Date: Tue, 5 Nov 2024 20:19:53 +0800 Subject: [PATCH 2/2] Update finetune_cli.py; formatting --- src/f5_tts/train/finetune_cli.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/f5_tts/train/finetune_cli.py b/src/f5_tts/train/finetune_cli.py index 13d861c..9e223aa 100644 --- a/src/f5_tts/train/finetune_cli.py +++ b/src/f5_tts/train/finetune_cli.py @@ -66,7 +66,7 @@ def parse_args(): "--bnb_optimizer", type=bool, default=False, - help="Use 8-bit Adam optimizer from bitsandbytes" + help="Use 8-bit Adam optimizer from bitsandbytes", ) return parser.parse_args()