mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
Improve argument handling (#209)
This commit is contained in:
+11
-17
@@ -455,7 +455,7 @@ def _add_model_arguments(parser: argparse.ArgumentParser) -> None:
|
|||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--pretrained_model_name_or_path",
|
"--pretrained_model_name_or_path",
|
||||||
type=str,
|
type=str,
|
||||||
default=None,
|
required=True,
|
||||||
help="Path to pretrained model or model identifier from huggingface.co/models.",
|
help="Path to pretrained model or model identifier from huggingface.co/models.",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
@@ -505,7 +505,7 @@ def _add_dataset_arguments(parser: argparse.ArgumentParser) -> None:
|
|||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--data_root",
|
"--data_root",
|
||||||
type=str,
|
type=str,
|
||||||
default=None,
|
required=True,
|
||||||
help=("A folder containing the training data."),
|
help=("A folder containing the training data."),
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
@@ -632,19 +632,19 @@ def _add_diffusion_arguments(parser: argparse.ArgumentParser) -> None:
|
|||||||
type=str,
|
type=str,
|
||||||
default="none",
|
default="none",
|
||||||
choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "none"],
|
choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "none"],
|
||||||
help=('We default to the "none" weighting scheme for uniform sampling and uniform loss'),
|
help='We default to the "none" weighting scheme for uniform sampling and uniform loss',
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--flow_logit_mean",
|
"--flow_logit_mean",
|
||||||
type=float,
|
type=float,
|
||||||
default=0.0,
|
default=0.0,
|
||||||
help="mean to use when using the `'logit_normal'` weighting scheme.",
|
help="Mean to use when using the `'logit_normal'` weighting scheme.",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--flow_logit_std",
|
"--flow_logit_std",
|
||||||
type=float,
|
type=float,
|
||||||
default=1.0,
|
default=1.0,
|
||||||
help="std to use when using the `'logit_normal'` weighting scheme.",
|
help="Standard deviation to use when using the `'logit_normal'` weighting scheme.",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--flow_mode_scale",
|
"--flow_mode_scale",
|
||||||
@@ -659,7 +659,7 @@ def _add_training_arguments(parser: argparse.ArgumentParser) -> None:
|
|||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--training_type",
|
"--training_type",
|
||||||
type=str,
|
type=str,
|
||||||
default=None,
|
required=True,
|
||||||
help="Type of training to perform. Choose between ['lora']",
|
help="Type of training to perform. Choose between ['lora']",
|
||||||
)
|
)
|
||||||
parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.")
|
parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.")
|
||||||
@@ -676,10 +676,10 @@ def _add_training_arguments(parser: argparse.ArgumentParser) -> None:
|
|||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--batch_size",
|
"--batch_size",
|
||||||
type=int,
|
type=int,
|
||||||
default=4,
|
default=1,
|
||||||
help="Batch size (per device) for the training dataloader.",
|
help="Batch size (per device) for the training dataloader.",
|
||||||
)
|
)
|
||||||
parser.add_argument("--train_epochs", type=int, default=1)
|
parser.add_argument("--train_epochs", type=int, default=1, help="Number of training epochs.")
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--train_steps",
|
"--train_steps",
|
||||||
type=int,
|
type=int,
|
||||||
@@ -735,13 +735,11 @@ def _add_training_arguments(parser: argparse.ArgumentParser) -> None:
|
|||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--enable_slicing",
|
"--enable_slicing",
|
||||||
action="store_true",
|
action="store_true",
|
||||||
default=False,
|
|
||||||
help="Whether or not to use VAE slicing for saving memory.",
|
help="Whether or not to use VAE slicing for saving memory.",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--enable_tiling",
|
"--enable_tiling",
|
||||||
action="store_true",
|
action="store_true",
|
||||||
default=False,
|
|
||||||
help="Whether or not to use VAE tiling for saving memory.",
|
help="Whether or not to use VAE tiling for saving memory.",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -756,7 +754,6 @@ def _add_optimizer_arguments(parser: argparse.ArgumentParser) -> None:
|
|||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--scale_lr",
|
"--scale_lr",
|
||||||
action="store_true",
|
action="store_true",
|
||||||
default=False,
|
|
||||||
help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
|
help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
@@ -877,7 +874,6 @@ def _add_validation_arguments(parser: argparse.ArgumentParser) -> None:
|
|||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--enable_model_cpu_offload",
|
"--enable_model_cpu_offload",
|
||||||
action="store_true",
|
action="store_true",
|
||||||
default=False,
|
|
||||||
help="Whether or not to enable model-wise CPU offloading when performing validation/testing to save memory.",
|
help="Whether or not to enable model-wise CPU offloading when performing validation/testing to save memory.",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -904,7 +900,7 @@ def _add_miscellaneous_arguments(parser: argparse.ArgumentParser) -> None:
|
|||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--output_dir",
|
"--output_dir",
|
||||||
type=str,
|
type=str,
|
||||||
default="finetrainer-training",
|
default="finetrainers-training",
|
||||||
help="The output directory where the model predictions and checkpoints will be written.",
|
help="The output directory where the model predictions and checkpoints will be written.",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
@@ -931,10 +927,8 @@ def _add_miscellaneous_arguments(parser: argparse.ArgumentParser) -> None:
|
|||||||
"--report_to",
|
"--report_to",
|
||||||
type=str,
|
type=str,
|
||||||
default="none",
|
default="none",
|
||||||
help=(
|
choices=["none", "wandb"],
|
||||||
'The integration to report the results and logs to. Supported platforms are `"tensorboard"`'
|
help="The integration to report the results and logs to.",
|
||||||
' (default), `"wandb"` and `"comet_ml"`. Use `"all"` to report to all integrations.'
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user