minor fixes

This commit is contained in:
Sourab Mangrulkar
2023-01-04 15:16:32 +05:30
parent 4e577cf6e3
commit c8d3def60e
2 changed files with 16 additions and 2 deletions
+2 -2
View File
@@ -67,14 +67,14 @@ Hardware: Single A100 80GB GPU with CPU RAM above 64G
| Model | Full Finetuning | PET-LoRA |
| --------- | ---- | ---- |
| CompVis/stable-diffusion-v1-4 | 66.4GB GPU / 3.97GB CPU / 22 minutes | 15.5GB GPU / 3.84GB CPU / 5.5 minutes |
| CompVis/stable-diffusion-v1-4 | 27.5GB GPU / 3.97GB CPU / 22 minutes | 15.5GB GPU / 3.84GB CPU / 5.5 minutes |
**Training**
An example of using LoRA for parameter efficient dreambooth training is given in `~examples/lora_dreambooth/train_dreambooth.py`
```bash
export MODEL_NAME="stabilityai/stable-diffusion-2-1" #"CompVis/stable-diffusion-v1-4"
export MODEL_NAME= "CompVis/stable-diffusion-v1-4" #"stabilityai/stable-diffusion-2-1"
export INSTANCE_DIR="path-to-instance-images"
export CLASS_DIR="path-to-class-images"
export OUTPUT_DIR="path-to-save-model"
@@ -155,6 +155,12 @@ def parse_args(input_args=None):
parser.add_argument("--lora_r", type=int, default=8, help="LoRA rank, only used if use_lora is True")
parser.add_argument("--lora_alpha", type=int, default=32, help="LoRA alpha, only used if use_lora is True")
parser.add_argument("--lora_dropout", type=float, default=0.0, help="LoRA dropout, only used if use_lora is True")
parser.add_argument(
"--lora_bias",
type=str,
default="none",
help="Bias type for LoRA. Can be 'none', 'all' or 'lora_only', only used if use_lora is True",
)
parser.add_argument(
"--lora_text_encoder_r",
type=int,
@@ -173,6 +179,12 @@ def parse_args(input_args=None):
default=0.0,
help="LoRA dropout for text encoder, only used if `use_lora` and `train_text_encoder` are True",
)
parser.add_argument(
"--lora_text_encoder_bias",
type=str,
default="none",
help="Bias type for LoRA. Can be 'none', 'all' or 'lora_only', only used if use_lora and `train_text_encoder` are True",
)
parser.add_argument(
"--train_batch_size", type=int, default=4, help="Batch size (per device) for the training dataloader."
@@ -675,6 +687,7 @@ def main(args):
lora_alpha=args.lora_alpha,
target_modules=UNET_TARGET_MODULES,
lora_dropout=args.lora_dropout,
bias=args.lora_bias,
)
unet = LoRAModel(config, unet)
print_trainable_parameters(unet)
@@ -689,6 +702,7 @@ def main(args):
lora_alpha=args.lora_text_encoder_alpha,
target_modules=TEXT_ENCODER_TARGET_MODULES,
lora_dropout=args.lora_text_encoder_dropout,
bias=args.lora_text_encoder_bias,
)
text_encoder = LoRAModel(config, text_encoder)
print_trainable_parameters(text_encoder)