Torba2345 commited on
Commit
8581949
·
verified ·
1 Parent(s): f9fa4a2

Upload train_dreambooth.py

Browse files
Files changed (1) hide show
  1. train_dreambooth.py +1444 -0
train_dreambooth.py ADDED
@@ -0,0 +1,1444 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python
2
+ # coding=utf-8
3
+ # Copyright 2025 The HuggingFace Inc. team. All rights reserved.
4
+ #
5
+ # Licensed under the Apache License, Version 2.0 (the "License");
6
+ # you may not use this file except in compliance with the License.
7
+ # You may obtain a copy of the License at
8
+ #
9
+ # http://www.apache.org/licenses/LICENSE-2.0
10
+ #
11
+ # Unless required by applicable law or agreed to in writing, software
12
+ # distributed under the License is distributed on an "AS IS" BASIS,
13
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
14
+ # See the License for the specific language governing permissions and
15
+
16
+ import argparse
17
+ import copy
18
+ import gc
19
+ import importlib
20
+ import itertools
21
+ import logging
22
+ import math
23
+ import os
24
+ import shutil
25
+ import warnings
26
+ from pathlib import Path
27
+
28
+ import numpy as np
29
+ import torch
30
+ import torch.nn.functional as F
31
+ import torch.utils.checkpoint
32
+ import transformers
33
+ from accelerate import Accelerator
34
+ from accelerate.logging import get_logger
35
+ from accelerate.utils import ProjectConfiguration, set_seed
36
+ from huggingface_hub import create_repo, model_info, upload_folder
37
+ from huggingface_hub.utils import insecure_hashlib
38
+ from packaging import version
39
+ from PIL import Image
40
+ from PIL.ImageOps import exif_transpose
41
+ from torch.utils.data import Dataset
42
+ from torchvision import transforms
43
+ from tqdm.auto import tqdm
44
+ from transformers import AutoTokenizer, PretrainedConfig
45
+
46
+ import diffusers
47
+ from diffusers import (
48
+ AutoencoderKL,
49
+ DDPMScheduler,
50
+ DiffusionPipeline,
51
+ StableDiffusionPipeline,
52
+ UNet2DConditionModel,
53
+ )
54
+ from diffusers.optimization import get_scheduler
55
+ from diffusers.training_utils import compute_snr
56
+ from diffusers.utils import check_min_version, is_wandb_available
57
+ from diffusers.utils.hub_utils import load_or_create_model_card, populate_model_card
58
+ from diffusers.utils.import_utils import is_xformers_available
59
+ from diffusers.utils.torch_utils import is_compiled_module
60
+
61
+
62
+ if is_wandb_available():
63
+ import wandb
64
+
65
+ # Will error if the minimal version of diffusers is not installed. Remove at your own risks.
66
+ check_min_version("0.34.0.dev0")
67
+
68
+ logger = get_logger(__name__)
69
+
70
+
71
+ def save_model_card(
72
+ repo_id: str,
73
+ images: list = None,
74
+ base_model: str = None,
75
+ train_text_encoder=False,
76
+ prompt: str = None,
77
+ repo_folder: str = None,
78
+ pipeline: DiffusionPipeline = None,
79
+ ):
80
+ img_str = ""
81
+ if images is not None:
82
+ for i, image in enumerate(images):
83
+ image.save(os.path.join(repo_folder, f"image_{i}.png"))
84
+ img_str += f"![img_{i}](./image_{i}.png)\n"
85
+
86
+ model_description = f"""
87
+ # DreamBooth - {repo_id}
88
+
89
+ This is a dreambooth model derived from {base_model}. The weights were trained on {prompt} using [DreamBooth](https://dreambooth.github.io/).
90
+ You can find some example images in the following. \n
91
+ {img_str}
92
+
93
+ DreamBooth for the text encoder was enabled: {train_text_encoder}.
94
+ """
95
+ model_card = load_or_create_model_card(
96
+ repo_id_or_path=repo_id,
97
+ from_training=True,
98
+ license="creativeml-openrail-m",
99
+ base_model=base_model,
100
+ prompt=prompt,
101
+ model_description=model_description,
102
+ inference=True,
103
+ )
104
+
105
+ tags = ["text-to-image", "dreambooth", "diffusers-training"]
106
+ if isinstance(pipeline, StableDiffusionPipeline):
107
+ tags.extend(["stable-diffusion", "stable-diffusion-diffusers"])
108
+ else:
109
+ tags.extend(["if", "if-diffusers"])
110
+ model_card = populate_model_card(model_card, tags=tags)
111
+
112
+ model_card.save(os.path.join(repo_folder, "README.md"))
113
+
114
+
115
+ def log_validation(
116
+ text_encoder,
117
+ tokenizer,
118
+ unet,
119
+ vae,
120
+ args,
121
+ accelerator,
122
+ weight_dtype,
123
+ global_step,
124
+ prompt_embeds,
125
+ negative_prompt_embeds,
126
+ ):
127
+ logger.info(
128
+ f"Running validation... \n Generating {args.num_validation_images} images with prompt:"
129
+ f" {args.validation_prompt}."
130
+ )
131
+
132
+ pipeline_args = {}
133
+
134
+ if vae is not None:
135
+ pipeline_args["vae"] = vae
136
+
137
+ # create pipeline (note: unet and vae are loaded again in float32)
138
+ pipeline = DiffusionPipeline.from_pretrained(
139
+ args.pretrained_model_name_or_path,
140
+ tokenizer=tokenizer,
141
+ text_encoder=text_encoder,
142
+ unet=unet,
143
+ revision=args.revision,
144
+ variant=args.variant,
145
+ torch_dtype=weight_dtype,
146
+ **pipeline_args,
147
+ )
148
+
149
+ # We train on the simplified learning objective. If we were previously predicting a variance, we need the scheduler to ignore it
150
+ scheduler_args = {}
151
+
152
+ if "variance_type" in pipeline.scheduler.config:
153
+ variance_type = pipeline.scheduler.config.variance_type
154
+
155
+ if variance_type in ["learned", "learned_range"]:
156
+ variance_type = "fixed_small"
157
+
158
+ scheduler_args["variance_type"] = variance_type
159
+
160
+ module = importlib.import_module("diffusers")
161
+ scheduler_class = getattr(module, args.validation_scheduler)
162
+ pipeline.scheduler = scheduler_class.from_config(pipeline.scheduler.config, **scheduler_args)
163
+ pipeline = pipeline.to(accelerator.device)
164
+ pipeline.set_progress_bar_config(disable=True)
165
+
166
+ if args.pre_compute_text_embeddings:
167
+ pipeline_args = {
168
+ "prompt_embeds": prompt_embeds,
169
+ "negative_prompt_embeds": negative_prompt_embeds,
170
+ }
171
+ else:
172
+ pipeline_args = {"prompt": args.validation_prompt}
173
+
174
+ # run inference
175
+ generator = None if args.seed is None else torch.Generator(device=accelerator.device).manual_seed(args.seed)
176
+ images = []
177
+ if args.validation_images is None:
178
+ for _ in range(args.num_validation_images):
179
+ with torch.autocast("cuda"):
180
+ image = pipeline(**pipeline_args, num_inference_steps=25, generator=generator).images[0]
181
+ images.append(image)
182
+ else:
183
+ for image in args.validation_images:
184
+ image = Image.open(image)
185
+ image = pipeline(**pipeline_args, image=image, generator=generator).images[0]
186
+ images.append(image)
187
+
188
+ for tracker in accelerator.trackers:
189
+ if tracker.name == "tensorboard":
190
+ np_images = np.stack([np.asarray(img) for img in images])
191
+ tracker.writer.add_images("validation", np_images, global_step, dataformats="NHWC")
192
+ if tracker.name == "wandb":
193
+ tracker.log(
194
+ {
195
+ "validation": [
196
+ wandb.Image(image, caption=f"{i}: {args.validation_prompt}") for i, image in enumerate(images)
197
+ ]
198
+ }
199
+ )
200
+
201
+ del pipeline
202
+ torch.cuda.empty_cache()
203
+
204
+ return images
205
+
206
+
207
+ def import_model_class_from_model_name_or_path(pretrained_model_name_or_path: str, revision: str):
208
+ text_encoder_config = PretrainedConfig.from_pretrained(
209
+ pretrained_model_name_or_path,
210
+ subfolder="text_encoder",
211
+ revision=revision,
212
+ )
213
+ model_class = text_encoder_config.architectures[0]
214
+
215
+ if model_class == "CLIPTextModel":
216
+ from transformers import CLIPTextModel
217
+
218
+ return CLIPTextModel
219
+ elif model_class == "RobertaSeriesModelWithTransformation":
220
+ from diffusers.pipelines.alt_diffusion.modeling_roberta_series import RobertaSeriesModelWithTransformation
221
+
222
+ return RobertaSeriesModelWithTransformation
223
+ elif model_class == "T5EncoderModel":
224
+ from transformers import T5EncoderModel
225
+
226
+ return T5EncoderModel
227
+ else:
228
+ raise ValueError(f"{model_class} is not supported.")
229
+
230
+
231
+ def parse_args(input_args=None):
232
+ parser = argparse.ArgumentParser(description="Simple example of a training script.")
233
+ parser.add_argument(
234
+ "--pretrained_model_name_or_path",
235
+ type=str,
236
+ default=None,
237
+ required=True,
238
+ help="Path to pretrained model or model identifier from huggingface.co/models.",
239
+ )
240
+ parser.add_argument(
241
+ "--revision",
242
+ type=str,
243
+ default=None,
244
+ required=False,
245
+ help="Revision of pretrained model identifier from huggingface.co/models.",
246
+ )
247
+ parser.add_argument(
248
+ "--variant",
249
+ type=str,
250
+ default=None,
251
+ help="Variant of the model files of the pretrained model identifier from huggingface.co/models, 'e.g.' fp16",
252
+ )
253
+ parser.add_argument(
254
+ "--tokenizer_name",
255
+ type=str,
256
+ default=None,
257
+ help="Pretrained tokenizer name or path if not the same as model_name",
258
+ )
259
+ parser.add_argument(
260
+ "--instance_data_dir",
261
+ type=str,
262
+ default=None,
263
+ required=True,
264
+ help="A folder containing the training data of instance images.",
265
+ )
266
+ parser.add_argument(
267
+ "--class_data_dir",
268
+ type=str,
269
+ default=None,
270
+ required=False,
271
+ help="A folder containing the training data of class images.",
272
+ )
273
+ parser.add_argument(
274
+ "--instance_prompt",
275
+ type=str,
276
+ default=None,
277
+ required=True,
278
+ help="The prompt with identifier specifying the instance",
279
+ )
280
+ parser.add_argument(
281
+ "--class_prompt",
282
+ type=str,
283
+ default=None,
284
+ help="The prompt to specify images in the same class as provided instance images.",
285
+ )
286
+ parser.add_argument(
287
+ "--with_prior_preservation",
288
+ default=False,
289
+ action="store_true",
290
+ help="Flag to add prior preservation loss.",
291
+ )
292
+ parser.add_argument("--prior_loss_weight", type=float, default=1.0, help="The weight of prior preservation loss.")
293
+ parser.add_argument(
294
+ "--num_class_images",
295
+ type=int,
296
+ default=100,
297
+ help=(
298
+ "Minimal class images for prior preservation loss. If there are not enough images already present in"
299
+ " class_data_dir, additional images will be sampled with class_prompt."
300
+ ),
301
+ )
302
+ parser.add_argument(
303
+ "--output_dir",
304
+ type=str,
305
+ default="dreambooth-model",
306
+ help="The output directory where the model predictions and checkpoints will be written.",
307
+ )
308
+ parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.")
309
+ parser.add_argument(
310
+ "--resolution",
311
+ type=int,
312
+ default=512,
313
+ help=(
314
+ "The resolution for input images, all the images in the train/validation dataset will be resized to this"
315
+ " resolution"
316
+ ),
317
+ )
318
+ parser.add_argument(
319
+ "--center_crop",
320
+ default=False,
321
+ action="store_true",
322
+ help=(
323
+ "Whether to center crop the input images to the resolution. If not set, the images will be randomly"
324
+ " cropped. The images will be resized to the resolution first before cropping."
325
+ ),
326
+ )
327
+ parser.add_argument(
328
+ "--train_text_encoder",
329
+ action="store_true",
330
+ help="Whether to train the text encoder. If set, the text encoder should be float32 precision.",
331
+ )
332
+ parser.add_argument(
333
+ "--train_batch_size", type=int, default=4, help="Batch size (per device) for the training dataloader."
334
+ )
335
+ parser.add_argument(
336
+ "--sample_batch_size", type=int, default=4, help="Batch size (per device) for sampling images."
337
+ )
338
+ parser.add_argument("--num_train_epochs", type=int, default=1)
339
+ parser.add_argument(
340
+ "--max_train_steps",
341
+ type=int,
342
+ default=None,
343
+ help="Total number of training steps to perform. If provided, overrides num_train_epochs.",
344
+ )
345
+ parser.add_argument(
346
+ "--checkpointing_steps",
347
+ type=int,
348
+ default=500,
349
+ help=(
350
+ "Save a checkpoint of the training state every X updates. Checkpoints can be used for resuming training via `--resume_from_checkpoint`. "
351
+ "In the case that the checkpoint is better than the final trained model, the checkpoint can also be used for inference."
352
+ "Using a checkpoint for inference requires separate loading of the original pipeline and the individual checkpointed model components."
353
+ "See https://huggingface.co/docs/diffusers/main/en/training/dreambooth#performing-inference-using-a-saved-checkpoint for step by step"
354
+ "instructions."
355
+ ),
356
+ )
357
+ parser.add_argument(
358
+ "--checkpoints_total_limit",
359
+ type=int,
360
+ default=None,
361
+ help=(
362
+ "Max number of checkpoints to store. Passed as `total_limit` to the `Accelerator` `ProjectConfiguration`."
363
+ " See Accelerator::save_state https://huggingface.co/docs/accelerate/package_reference/accelerator#accelerate.Accelerator.save_state"
364
+ " for more details"
365
+ ),
366
+ )
367
+ parser.add_argument(
368
+ "--resume_from_checkpoint",
369
+ type=str,
370
+ default=None,
371
+ help=(
372
+ "Whether training should be resumed from a previous checkpoint. Use a path saved by"
373
+ ' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
374
+ ),
375
+ )
376
+ parser.add_argument(
377
+ "--gradient_accumulation_steps",
378
+ type=int,
379
+ default=1,
380
+ help="Number of updates steps to accumulate before performing a backward/update pass.",
381
+ )
382
+ parser.add_argument(
383
+ "--gradient_checkpointing",
384
+ action="store_true",
385
+ help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
386
+ )
387
+ parser.add_argument(
388
+ "--learning_rate",
389
+ type=float,
390
+ default=5e-6,
391
+ help="Initial learning rate (after the potential warmup period) to use.",
392
+ )
393
+ parser.add_argument(
394
+ "--scale_lr",
395
+ action="store_true",
396
+ default=False,
397
+ help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
398
+ )
399
+ parser.add_argument(
400
+ "--lr_scheduler",
401
+ type=str,
402
+ default="constant",
403
+ help=(
404
+ 'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
405
+ ' "constant", "constant_with_warmup"]'
406
+ ),
407
+ )
408
+ parser.add_argument(
409
+ "--lr_warmup_steps", type=int, default=500, help="Number of steps for the warmup in the lr scheduler."
410
+ )
411
+ parser.add_argument(
412
+ "--lr_num_cycles",
413
+ type=int,
414
+ default=1,
415
+ help="Number of hard resets of the lr in cosine_with_restarts scheduler.",
416
+ )
417
+ parser.add_argument("--lr_power", type=float, default=1.0, help="Power factor of the polynomial scheduler.")
418
+ parser.add_argument(
419
+ "--use_8bit_adam", action="store_true", help="Whether or not to use 8-bit Adam from bitsandbytes."
420
+ )
421
+ parser.add_argument(
422
+ "--dataloader_num_workers",
423
+ type=int,
424
+ default=0,
425
+ help=(
426
+ "Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process."
427
+ ),
428
+ )
429
+ parser.add_argument("--adam_beta1", type=float, default=0.9, help="The beta1 parameter for the Adam optimizer.")
430
+ parser.add_argument("--adam_beta2", type=float, default=0.999, help="The beta2 parameter for the Adam optimizer.")
431
+ parser.add_argument("--adam_weight_decay", type=float, default=1e-2, help="Weight decay to use.")
432
+ parser.add_argument("--adam_epsilon", type=float, default=1e-08, help="Epsilon value for the Adam optimizer")
433
+ parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.")
434
+ parser.add_argument("--push_to_hub", action="store_true", help="Whether or not to push the model to the Hub.")
435
+ parser.add_argument("--hub_token", type=str, default=None, help="The token to use to push to the Model Hub.")
436
+ parser.add_argument(
437
+ "--hub_model_id",
438
+ type=str,
439
+ default=None,
440
+ help="The name of the repository to keep in sync with the local `output_dir`.",
441
+ )
442
+ parser.add_argument(
443
+ "--logging_dir",
444
+ type=str,
445
+ default="logs",
446
+ help=(
447
+ "[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
448
+ " *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."
449
+ ),
450
+ )
451
+ parser.add_argument(
452
+ "--allow_tf32",
453
+ action="store_true",
454
+ help=(
455
+ "Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
456
+ " https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
457
+ ),
458
+ )
459
+ parser.add_argument(
460
+ "--report_to",
461
+ type=str,
462
+ default="tensorboard",
463
+ help=(
464
+ 'The integration to report the results and logs to. Supported platforms are `"tensorboard"`'
465
+ ' (default), `"wandb"` and `"comet_ml"`. Use `"all"` to report to all integrations.'
466
+ ),
467
+ )
468
+ parser.add_argument(
469
+ "--validation_prompt",
470
+ type=str,
471
+ default=None,
472
+ help="A prompt that is used during validation to verify that the model is learning.",
473
+ )
474
+ parser.add_argument(
475
+ "--num_validation_images",
476
+ type=int,
477
+ default=4,
478
+ help="Number of images that should be generated during validation with `validation_prompt`.",
479
+ )
480
+ parser.add_argument(
481
+ "--validation_steps",
482
+ type=int,
483
+ default=100,
484
+ help=(
485
+ "Run validation every X steps. Validation consists of running the prompt"
486
+ " `args.validation_prompt` multiple times: `args.num_validation_images`"
487
+ " and logging the images."
488
+ ),
489
+ )
490
+ parser.add_argument(
491
+ "--mixed_precision",
492
+ type=str,
493
+ default=None,
494
+ choices=["no", "fp16", "bf16"],
495
+ help=(
496
+ "Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
497
+ " 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
498
+ " flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."
499
+ ),
500
+ )
501
+ parser.add_argument(
502
+ "--prior_generation_precision",
503
+ type=str,
504
+ default=None,
505
+ choices=["no", "fp32", "fp16", "bf16"],
506
+ help=(
507
+ "Choose prior generation precision between fp32, fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
508
+ " 1.10.and an Nvidia Ampere GPU. Default to fp16 if a GPU is available else fp32."
509
+ ),
510
+ )
511
+ parser.add_argument("--local_rank", type=int, default=-1, help="For distributed training: local_rank")
512
+ parser.add_argument(
513
+ "--enable_xformers_memory_efficient_attention", action="store_true", help="Whether or not to use xformers."
514
+ )
515
+ parser.add_argument(
516
+ "--set_grads_to_none",
517
+ action="store_true",
518
+ help=(
519
+ "Save more memory by using setting grads to None instead of zero. Be aware, that this changes certain"
520
+ " behaviors, so disable this argument if it causes any problems. More info:"
521
+ " https://pytorch.org/docs/stable/generated/torch.optim.Optimizer.zero_grad.html"
522
+ ),
523
+ )
524
+
525
+ parser.add_argument(
526
+ "--offset_noise",
527
+ action="store_true",
528
+ default=False,
529
+ help=(
530
+ "Fine-tuning against a modified noise"
531
+ " See: https://www.crosslabs.org//blog/diffusion-with-offset-noise for more information."
532
+ ),
533
+ )
534
+ parser.add_argument(
535
+ "--snr_gamma",
536
+ type=float,
537
+ default=None,
538
+ help="SNR weighting gamma to be used if rebalancing the loss. Recommended value is 5.0. "
539
+ "More details here: https://arxiv.org/abs/2303.09556.",
540
+ )
541
+ parser.add_argument(
542
+ "--pre_compute_text_embeddings",
543
+ action="store_true",
544
+ help="Whether or not to pre-compute text embeddings. If text embeddings are pre-computed, the text encoder will not be kept in memory during training and will leave more GPU memory available for training the rest of the model. This is not compatible with `--train_text_encoder`.",
545
+ )
546
+ parser.add_argument(
547
+ "--tokenizer_max_length",
548
+ type=int,
549
+ default=None,
550
+ required=False,
551
+ help="The maximum length of the tokenizer. If not set, will default to the tokenizer's max length.",
552
+ )
553
+ parser.add_argument(
554
+ "--text_encoder_use_attention_mask",
555
+ action="store_true",
556
+ required=False,
557
+ help="Whether to use attention mask for the text encoder",
558
+ )
559
+ parser.add_argument(
560
+ "--skip_save_text_encoder", action="store_true", required=False, help="Set to not save text encoder"
561
+ )
562
+ parser.add_argument(
563
+ "--validation_images",
564
+ required=False,
565
+ default=None,
566
+ nargs="+",
567
+ help="Optional set of images to use for validation. Used when the target pipeline takes an initial image as input such as when training image variation or superresolution.",
568
+ )
569
+ parser.add_argument(
570
+ "--class_labels_conditioning",
571
+ required=False,
572
+ default=None,
573
+ help="The optional `class_label` conditioning to pass to the unet, available values are `timesteps`.",
574
+ )
575
+ parser.add_argument(
576
+ "--validation_scheduler",
577
+ type=str,
578
+ default="DPMSolverMultistepScheduler",
579
+ choices=["DPMSolverMultistepScheduler", "DDPMScheduler"],
580
+ help="Select which scheduler to use for validation. DDPMScheduler is recommended for DeepFloyd IF.",
581
+ )
582
+
583
+ if input_args is not None:
584
+ args = parser.parse_args(input_args)
585
+ else:
586
+ args = parser.parse_args()
587
+
588
+ env_local_rank = int(os.environ.get("LOCAL_RANK", -1))
589
+ if env_local_rank != -1 and env_local_rank != args.local_rank:
590
+ args.local_rank = env_local_rank
591
+
592
+ if args.with_prior_preservation:
593
+ if args.class_data_dir is None:
594
+ raise ValueError("You must specify a data directory for class images.")
595
+ if args.class_prompt is None:
596
+ raise ValueError("You must specify prompt for class images.")
597
+ else:
598
+ # logger is not available yet
599
+ if args.class_data_dir is not None:
600
+ warnings.warn("You need not use --class_data_dir without --with_prior_preservation.")
601
+ if args.class_prompt is not None:
602
+ warnings.warn("You need not use --class_prompt without --with_prior_preservation.")
603
+
604
+ if args.train_text_encoder and args.pre_compute_text_embeddings:
605
+ raise ValueError("`--train_text_encoder` cannot be used with `--pre_compute_text_embeddings`")
606
+
607
+ return args
608
+
609
+
610
+ class DreamBoothDataset(Dataset):
611
+ """
612
+ A dataset to prepare the instance and class images with the prompts for fine-tuning the model.
613
+ It pre-processes the images and the tokenizes prompts.
614
+ """
615
+
616
+ def __init__(
617
+ self,
618
+ instance_data_root,
619
+ instance_prompt,
620
+ tokenizer,
621
+ class_data_root=None,
622
+ class_prompt=None,
623
+ class_num=None,
624
+ size=512,
625
+ center_crop=False,
626
+ encoder_hidden_states=None,
627
+ class_prompt_encoder_hidden_states=None,
628
+ tokenizer_max_length=None,
629
+ ):
630
+ self.size = size
631
+ self.center_crop = center_crop
632
+ self.tokenizer = tokenizer
633
+ self.encoder_hidden_states = encoder_hidden_states
634
+ self.class_prompt_encoder_hidden_states = class_prompt_encoder_hidden_states
635
+ self.tokenizer_max_length = tokenizer_max_length
636
+
637
+ self.instance_data_root = Path(instance_data_root)
638
+ if not self.instance_data_root.exists():
639
+ raise ValueError(f"Instance {self.instance_data_root} images root doesn't exists.")
640
+
641
+ self.instance_images_path = list(Path(instance_data_root).iterdir())
642
+ self.num_instance_images = len(self.instance_images_path)
643
+ self.instance_prompt = instance_prompt
644
+ self._length = self.num_instance_images
645
+
646
+ if class_data_root is not None:
647
+ self.class_data_root = Path(class_data_root)
648
+ self.class_data_root.mkdir(parents=True, exist_ok=True)
649
+ self.class_images_path = list(self.class_data_root.iterdir())
650
+ if class_num is not None:
651
+ self.num_class_images = min(len(self.class_images_path), class_num)
652
+ else:
653
+ self.num_class_images = len(self.class_images_path)
654
+ self._length = max(self.num_class_images, self.num_instance_images)
655
+ self.class_prompt = class_prompt
656
+ else:
657
+ self.class_data_root = None
658
+
659
+ self.image_transforms = transforms.Compose(
660
+ [
661
+ transforms.Resize(size, interpolation=transforms.InterpolationMode.BILINEAR),
662
+ transforms.CenterCrop(size) if center_crop else transforms.RandomCrop(size),
663
+ transforms.ToTensor(),
664
+ transforms.Normalize([0.5], [0.5]),
665
+ ]
666
+ )
667
+
668
+ def __len__(self):
669
+ return self._length
670
+
671
+ def __getitem__(self, index):
672
+ example = {}
673
+ instance_image = Image.open(self.instance_images_path[index % self.num_instance_images])
674
+ instance_image = exif_transpose(instance_image)
675
+
676
+ if not instance_image.mode == "RGB":
677
+ instance_image = instance_image.convert("RGB")
678
+ example["instance_images"] = self.image_transforms(instance_image)
679
+
680
+ if self.encoder_hidden_states is not None:
681
+ example["instance_prompt_ids"] = self.encoder_hidden_states
682
+ else:
683
+ text_inputs = tokenize_prompt(
684
+ self.tokenizer, self.instance_prompt, tokenizer_max_length=self.tokenizer_max_length
685
+ )
686
+ example["instance_prompt_ids"] = text_inputs.input_ids
687
+ example["instance_attention_mask"] = text_inputs.attention_mask
688
+
689
+ if self.class_data_root:
690
+ class_image = Image.open(self.class_images_path[index % self.num_class_images])
691
+ class_image = exif_transpose(class_image)
692
+
693
+ if not class_image.mode == "RGB":
694
+ class_image = class_image.convert("RGB")
695
+ example["class_images"] = self.image_transforms(class_image)
696
+
697
+ if self.class_prompt_encoder_hidden_states is not None:
698
+ example["class_prompt_ids"] = self.class_prompt_encoder_hidden_states
699
+ else:
700
+ class_text_inputs = tokenize_prompt(
701
+ self.tokenizer, self.class_prompt, tokenizer_max_length=self.tokenizer_max_length
702
+ )
703
+ example["class_prompt_ids"] = class_text_inputs.input_ids
704
+ example["class_attention_mask"] = class_text_inputs.attention_mask
705
+
706
+ return example
707
+
708
+
709
+ def collate_fn(examples, with_prior_preservation=False):
710
+ has_attention_mask = "instance_attention_mask" in examples[0]
711
+
712
+ input_ids = [example["instance_prompt_ids"] for example in examples]
713
+ pixel_values = [example["instance_images"] for example in examples]
714
+
715
+ if has_attention_mask:
716
+ attention_mask = [example["instance_attention_mask"] for example in examples]
717
+
718
+ # Concat class and instance examples for prior preservation.
719
+ # We do this to avoid doing two forward passes.
720
+ if with_prior_preservation:
721
+ input_ids += [example["class_prompt_ids"] for example in examples]
722
+ pixel_values += [example["class_images"] for example in examples]
723
+
724
+ if has_attention_mask:
725
+ attention_mask += [example["class_attention_mask"] for example in examples]
726
+
727
+ pixel_values = torch.stack(pixel_values)
728
+ pixel_values = pixel_values.to(memory_format=torch.contiguous_format).float()
729
+
730
+ input_ids = torch.cat(input_ids, dim=0)
731
+
732
+ batch = {
733
+ "input_ids": input_ids,
734
+ "pixel_values": pixel_values,
735
+ }
736
+
737
+ if has_attention_mask:
738
+ attention_mask = torch.cat(attention_mask, dim=0)
739
+ batch["attention_mask"] = attention_mask
740
+
741
+ return batch
742
+
743
+
744
+ class PromptDataset(Dataset):
745
+ """A simple dataset to prepare the prompts to generate class images on multiple GPUs."""
746
+
747
+ def __init__(self, prompt, num_samples):
748
+ self.prompt = prompt
749
+ self.num_samples = num_samples
750
+
751
+ def __len__(self):
752
+ return self.num_samples
753
+
754
+ def __getitem__(self, index):
755
+ example = {}
756
+ example["prompt"] = self.prompt
757
+ example["index"] = index
758
+ return example
759
+
760
+
761
+ def model_has_vae(args):
762
+ config_file_name = Path("vae", AutoencoderKL.config_name).as_posix()
763
+ if os.path.isdir(args.pretrained_model_name_or_path):
764
+ config_file_name = os.path.join(args.pretrained_model_name_or_path, config_file_name)
765
+ return os.path.isfile(config_file_name)
766
+ else:
767
+ files_in_repo = model_info(args.pretrained_model_name_or_path, revision=args.revision).siblings
768
+ return any(file.rfilename == config_file_name for file in files_in_repo)
769
+
770
+
771
+ def tokenize_prompt(tokenizer, prompt, tokenizer_max_length=None):
772
+ if tokenizer_max_length is not None:
773
+ max_length = tokenizer_max_length
774
+ else:
775
+ max_length = tokenizer.model_max_length
776
+
777
+ text_inputs = tokenizer(
778
+ prompt,
779
+ truncation=True,
780
+ padding="max_length",
781
+ max_length=max_length,
782
+ return_tensors="pt",
783
+ )
784
+
785
+ return text_inputs
786
+
787
+
788
+ def encode_prompt(text_encoder, input_ids, attention_mask, text_encoder_use_attention_mask=None):
789
+ text_input_ids = input_ids.to(text_encoder.device)
790
+
791
+ if text_encoder_use_attention_mask:
792
+ attention_mask = attention_mask.to(text_encoder.device)
793
+ else:
794
+ attention_mask = None
795
+
796
+ prompt_embeds = text_encoder(
797
+ text_input_ids,
798
+ attention_mask=attention_mask,
799
+ return_dict=False,
800
+ )
801
+ prompt_embeds = prompt_embeds[0]
802
+
803
+ return prompt_embeds
804
+
805
+
806
+ def main(args):
807
+ if args.report_to == "wandb" and args.hub_token is not None:
808
+ raise ValueError(
809
+ "You cannot use both --report_to=wandb and --hub_token due to a security risk of exposing your token."
810
+ " Please use `huggingface-cli login` to authenticate with the Hub."
811
+ )
812
+
813
+ logging_dir = Path(args.output_dir, args.logging_dir)
814
+
815
+ accelerator_project_config = ProjectConfiguration(project_dir=args.output_dir, logging_dir=logging_dir)
816
+
817
+ accelerator = Accelerator(
818
+ gradient_accumulation_steps=args.gradient_accumulation_steps,
819
+ mixed_precision=args.mixed_precision,
820
+ log_with=args.report_to,
821
+ project_config=accelerator_project_config,
822
+ )
823
+
824
+ # Disable AMP for MPS.
825
+ if torch.backends.mps.is_available():
826
+ accelerator.native_amp = False
827
+
828
+ if args.report_to == "wandb":
829
+ if not is_wandb_available():
830
+ raise ImportError("Make sure to install wandb if you want to use it for logging during training.")
831
+
832
+ # Currently, it's not possible to do gradient accumulation when training two models with accelerate.accumulate
833
+ # This will be enabled soon in accelerate. For now, we don't allow gradient accumulation when training two models.
834
+ # TODO (patil-suraj): Remove this check when gradient accumulation with two models is enabled in accelerate.
835
+ if args.train_text_encoder and args.gradient_accumulation_steps > 1 and accelerator.num_processes > 1:
836
+ raise ValueError(
837
+ "Gradient accumulation is not supported when training the text encoder in distributed training. "
838
+ "Please set gradient_accumulation_steps to 1. This feature will be supported in the future."
839
+ )
840
+
841
+ # Make one log on every process with the configuration for debugging.
842
+ logging.basicConfig(
843
+ format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",
844
+ datefmt="%m/%d/%Y %H:%M:%S",
845
+ level=logging.INFO,
846
+ )
847
+ logger.info(accelerator.state, main_process_only=False)
848
+ if accelerator.is_local_main_process:
849
+ transformers.utils.logging.set_verbosity_warning()
850
+ diffusers.utils.logging.set_verbosity_info()
851
+ else:
852
+ transformers.utils.logging.set_verbosity_error()
853
+ diffusers.utils.logging.set_verbosity_error()
854
+
855
+ # If passed along, set the training seed now.
856
+ if args.seed is not None:
857
+ set_seed(args.seed)
858
+
859
+ # Generate class images if prior preservation is enabled.
860
+ if args.with_prior_preservation:
861
+ class_images_dir = Path(args.class_data_dir)
862
+ if not class_images_dir.exists():
863
+ class_images_dir.mkdir(parents=True)
864
+ cur_class_images = len(list(class_images_dir.iterdir()))
865
+
866
+ if cur_class_images < args.num_class_images:
867
+ torch_dtype = torch.float16 if accelerator.device.type == "cuda" else torch.float32
868
+ if args.prior_generation_precision == "fp32":
869
+ torch_dtype = torch.float32
870
+ elif args.prior_generation_precision == "fp16":
871
+ torch_dtype = torch.float16
872
+ elif args.prior_generation_precision == "bf16":
873
+ torch_dtype = torch.bfloat16
874
+ pipeline = DiffusionPipeline.from_pretrained(
875
+ args.pretrained_model_name_or_path,
876
+ torch_dtype=torch_dtype,
877
+ safety_checker=None,
878
+ revision=args.revision,
879
+ variant=args.variant,
880
+ )
881
+ pipeline.set_progress_bar_config(disable=True)
882
+
883
+ num_new_images = args.num_class_images - cur_class_images
884
+ logger.info(f"Number of class images to sample: {num_new_images}.")
885
+
886
+ sample_dataset = PromptDataset(args.class_prompt, num_new_images)
887
+ sample_dataloader = torch.utils.data.DataLoader(sample_dataset, batch_size=args.sample_batch_size)
888
+
889
+ sample_dataloader = accelerator.prepare(sample_dataloader)
890
+ pipeline.to(accelerator.device)
891
+
892
+ for example in tqdm(
893
+ sample_dataloader, desc="Generating class images", disable=not accelerator.is_local_main_process
894
+ ):
895
+ images = pipeline(example["prompt"]).images
896
+
897
+ for i, image in enumerate(images):
898
+ hash_image = insecure_hashlib.sha1(image.tobytes()).hexdigest()
899
+ image_filename = class_images_dir / f"{example['index'][i] + cur_class_images}-{hash_image}.jpg"
900
+ image.save(image_filename)
901
+
902
+ del pipeline
903
+ if torch.cuda.is_available():
904
+ torch.cuda.empty_cache()
905
+
906
+ # Handle the repository creation
907
+ if accelerator.is_main_process:
908
+ if args.output_dir is not None:
909
+ os.makedirs(args.output_dir, exist_ok=True)
910
+
911
+ if args.push_to_hub:
912
+ repo_id = create_repo(
913
+ repo_id=args.hub_model_id or Path(args.output_dir).name, exist_ok=True, token=args.hub_token
914
+ ).repo_id
915
+
916
+ # Load the tokenizer
917
+ if args.tokenizer_name:
918
+ tokenizer = AutoTokenizer.from_pretrained(args.tokenizer_name, revision=args.revision, use_fast=False)
919
+ elif args.pretrained_model_name_or_path:
920
+ tokenizer = AutoTokenizer.from_pretrained(
921
+ args.pretrained_model_name_or_path,
922
+ subfolder="tokenizer",
923
+ revision=args.revision,
924
+ use_fast=False,
925
+ )
926
+
927
+ # import correct text encoder class
928
+ text_encoder_cls = import_model_class_from_model_name_or_path(args.pretrained_model_name_or_path, args.revision)
929
+
930
+ # Load scheduler and models
931
+ noise_scheduler = DDPMScheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="scheduler")
932
+ text_encoder = text_encoder_cls.from_pretrained(
933
+ args.pretrained_model_name_or_path, subfolder="text_encoder", revision=args.revision, variant=args.variant
934
+ )
935
+
936
+ if model_has_vae(args):
937
+ vae = AutoencoderKL.from_pretrained(
938
+ args.pretrained_model_name_or_path, subfolder="vae", revision=args.revision, variant=args.variant
939
+ )
940
+ else:
941
+ vae = None
942
+
943
+ unet = UNet2DConditionModel.from_pretrained(
944
+ args.pretrained_model_name_or_path, subfolder="unet", revision=args.revision, variant=args.variant
945
+ )
946
+
947
+ def unwrap_model(model):
948
+ model = accelerator.unwrap_model(model)
949
+ model = model._orig_mod if is_compiled_module(model) else model
950
+ return model
951
+
952
+ # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format
953
+ def save_model_hook(models, weights, output_dir):
954
+ if accelerator.is_main_process:
955
+ for model in models:
956
+ sub_dir = "unet" if isinstance(model, type(unwrap_model(unet))) else "text_encoder"
957
+ model.save_pretrained(os.path.join(output_dir, sub_dir))
958
+
959
+ # make sure to pop weight so that corresponding model is not saved again
960
+ weights.pop()
961
+
962
+ def load_model_hook(models, input_dir):
963
+ while len(models) > 0:
964
+ # pop models so that they are not loaded again
965
+ model = models.pop()
966
+
967
+ if isinstance(model, type(unwrap_model(text_encoder))):
968
+ # load transformers style into model
969
+ load_model = text_encoder_cls.from_pretrained(input_dir, subfolder="text_encoder")
970
+ model.config = load_model.config
971
+ else:
972
+ # load diffusers style into model
973
+ load_model = UNet2DConditionModel.from_pretrained(input_dir, subfolder="unet")
974
+ model.register_to_config(**load_model.config)
975
+
976
+ model.load_state_dict(load_model.state_dict())
977
+ del load_model
978
+
979
+ accelerator.register_save_state_pre_hook(save_model_hook)
980
+ accelerator.register_load_state_pre_hook(load_model_hook)
981
+
982
+ if vae is not None:
983
+ vae.requires_grad_(False)
984
+
985
+ if not args.train_text_encoder:
986
+ text_encoder.requires_grad_(False)
987
+
988
+ if args.enable_xformers_memory_efficient_attention:
989
+ if is_xformers_available():
990
+ import xformers
991
+
992
+ xformers_version = version.parse(xformers.__version__)
993
+ if xformers_version == version.parse("0.0.16"):
994
+ logger.warning(
995
+ "xFormers 0.0.16 cannot be used for training in some GPUs. If you observe problems during training, please update xFormers to at least 0.0.17. See https://huggingface.co/docs/diffusers/main/en/optimization/xformers for more details."
996
+ )
997
+ unet.enable_xformers_memory_efficient_attention()
998
+ else:
999
+ raise ValueError("xformers is not available. Make sure it is installed correctly")
1000
+
1001
+ if args.gradient_checkpointing:
1002
+ unet.enable_gradient_checkpointing()
1003
+ if args.train_text_encoder:
1004
+ text_encoder.gradient_checkpointing_enable()
1005
+
1006
+ # Check that all trainable models are in full precision
1007
+ low_precision_error_string = (
1008
+ "Please make sure to always have all model weights in full float32 precision when starting training - even if"
1009
+ " doing mixed precision training. copy of the weights should still be float32."
1010
+ )
1011
+
1012
+ if unwrap_model(unet).dtype != torch.float32:
1013
+ raise ValueError(f"Unet loaded as datatype {unwrap_model(unet).dtype}. {low_precision_error_string}")
1014
+
1015
+ if args.train_text_encoder and unwrap_model(text_encoder).dtype != torch.float32:
1016
+ raise ValueError(
1017
+ f"Text encoder loaded as datatype {unwrap_model(text_encoder).dtype}. {low_precision_error_string}"
1018
+ )
1019
+
1020
+ # Enable TF32 for faster training on Ampere GPUs,
1021
+ # cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices
1022
+ if args.allow_tf32:
1023
+ torch.backends.cuda.matmul.allow_tf32 = True
1024
+
1025
+ if args.scale_lr:
1026
+ args.learning_rate = (
1027
+ args.learning_rate * args.gradient_accumulation_steps * args.train_batch_size * accelerator.num_processes
1028
+ )
1029
+
1030
+ # Use 8-bit Adam for lower memory usage or to fine-tune the model in 16GB GPUs
1031
+ if args.use_8bit_adam:
1032
+ try:
1033
+ import bitsandbytes as bnb
1034
+ except ImportError:
1035
+ raise ImportError(
1036
+ "To use 8-bit Adam, please install the bitsandbytes library: `pip install bitsandbytes`."
1037
+ )
1038
+
1039
+ optimizer_class = bnb.optim.AdamW8bit
1040
+ else:
1041
+ optimizer_class = torch.optim.AdamW
1042
+
1043
+ # Optimizer creation
1044
+ params_to_optimize = (
1045
+ itertools.chain(unet.parameters(), text_encoder.parameters()) if args.train_text_encoder else unet.parameters()
1046
+ )
1047
+ optimizer = optimizer_class(
1048
+ params_to_optimize,
1049
+ lr=args.learning_rate,
1050
+ betas=(args.adam_beta1, args.adam_beta2),
1051
+ weight_decay=args.adam_weight_decay,
1052
+ eps=args.adam_epsilon,
1053
+ )
1054
+
1055
+ if args.pre_compute_text_embeddings:
1056
+
1057
+ def compute_text_embeddings(prompt):
1058
+ with torch.no_grad():
1059
+ text_inputs = tokenize_prompt(tokenizer, prompt, tokenizer_max_length=args.tokenizer_max_length)
1060
+ prompt_embeds = encode_prompt(
1061
+ text_encoder,
1062
+ text_inputs.input_ids,
1063
+ text_inputs.attention_mask,
1064
+ text_encoder_use_attention_mask=args.text_encoder_use_attention_mask,
1065
+ )
1066
+
1067
+ return prompt_embeds
1068
+
1069
+ pre_computed_encoder_hidden_states = compute_text_embeddings(args.instance_prompt)
1070
+ validation_prompt_negative_prompt_embeds = compute_text_embeddings("")
1071
+
1072
+ if args.validation_prompt is not None:
1073
+ validation_prompt_encoder_hidden_states = compute_text_embeddings(args.validation_prompt)
1074
+ else:
1075
+ validation_prompt_encoder_hidden_states = None
1076
+
1077
+ if args.class_prompt is not None:
1078
+ pre_computed_class_prompt_encoder_hidden_states = compute_text_embeddings(args.class_prompt)
1079
+ else:
1080
+ pre_computed_class_prompt_encoder_hidden_states = None
1081
+
1082
+ text_encoder = None
1083
+ tokenizer = None
1084
+
1085
+ gc.collect()
1086
+ torch.cuda.empty_cache()
1087
+ else:
1088
+ pre_computed_encoder_hidden_states = None
1089
+ validation_prompt_encoder_hidden_states = None
1090
+ validation_prompt_negative_prompt_embeds = None
1091
+ pre_computed_class_prompt_encoder_hidden_states = None
1092
+
1093
+ # Dataset and DataLoaders creation:
1094
+ train_dataset = DreamBoothDataset(
1095
+ instance_data_root=args.instance_data_dir,
1096
+ instance_prompt=args.instance_prompt,
1097
+ class_data_root=args.class_data_dir if args.with_prior_preservation else None,
1098
+ class_prompt=args.class_prompt,
1099
+ class_num=args.num_class_images,
1100
+ tokenizer=tokenizer,
1101
+ size=args.resolution,
1102
+ center_crop=args.center_crop,
1103
+ encoder_hidden_states=pre_computed_encoder_hidden_states,
1104
+ class_prompt_encoder_hidden_states=pre_computed_class_prompt_encoder_hidden_states,
1105
+ tokenizer_max_length=args.tokenizer_max_length,
1106
+ )
1107
+
1108
+ train_dataloader = torch.utils.data.DataLoader(
1109
+ train_dataset,
1110
+ batch_size=args.train_batch_size,
1111
+ shuffle=True,
1112
+ collate_fn=lambda examples: collate_fn(examples, args.with_prior_preservation),
1113
+ num_workers=args.dataloader_num_workers,
1114
+ )
1115
+
1116
+ # Scheduler and math around the number of training steps.
1117
+ overrode_max_train_steps = False
1118
+ num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
1119
+ if args.max_train_steps is None:
1120
+ args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch
1121
+ overrode_max_train_steps = True
1122
+
1123
+ lr_scheduler = get_scheduler(
1124
+ args.lr_scheduler,
1125
+ optimizer=optimizer,
1126
+ num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes,
1127
+ num_training_steps=args.max_train_steps * accelerator.num_processes,
1128
+ num_cycles=args.lr_num_cycles,
1129
+ power=args.lr_power,
1130
+ )
1131
+
1132
+ # Prepare everything with our `accelerator`.
1133
+ if args.train_text_encoder:
1134
+ unet, text_encoder, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(
1135
+ unet, text_encoder, optimizer, train_dataloader, lr_scheduler
1136
+ )
1137
+ else:
1138
+ unet, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(
1139
+ unet, optimizer, train_dataloader, lr_scheduler
1140
+ )
1141
+
1142
+ # For mixed precision training we cast all non-trainable weights (vae, non-lora text_encoder and non-lora unet) to half-precision
1143
+ # as these weights are only used for inference, keeping weights in full precision is not required.
1144
+ weight_dtype = torch.float32
1145
+ if accelerator.mixed_precision == "fp16":
1146
+ weight_dtype = torch.float16
1147
+ elif accelerator.mixed_precision == "bf16":
1148
+ weight_dtype = torch.bfloat16
1149
+
1150
+ # Move vae and text_encoder to device and cast to weight_dtype
1151
+ if vae is not None:
1152
+ vae.to(accelerator.device, dtype=weight_dtype)
1153
+
1154
+ if not args.train_text_encoder and text_encoder is not None:
1155
+ text_encoder.to(accelerator.device, dtype=weight_dtype)
1156
+
1157
+ # We need to recalculate our total training steps as the size of the training dataloader may have changed.
1158
+ num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
1159
+ if overrode_max_train_steps:
1160
+ args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch
1161
+ # Afterwards we recalculate our number of training epochs
1162
+ args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
1163
+
1164
+ # We need to initialize the trackers we use, and also store our configuration.
1165
+ # The trackers initializes automatically on the main process.
1166
+ if accelerator.is_main_process:
1167
+ tracker_config = vars(copy.deepcopy(args))
1168
+ tracker_config.pop("validation_images")
1169
+ accelerator.init_trackers("dreambooth", config=tracker_config)
1170
+
1171
+ # Train!
1172
+ total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps
1173
+
1174
+ logger.info("***** Running training *****")
1175
+ logger.info(f" Num examples = {len(train_dataset)}")
1176
+ logger.info(f" Num batches each epoch = {len(train_dataloader)}")
1177
+ logger.info(f" Num Epochs = {args.num_train_epochs}")
1178
+ logger.info(f" Instantaneous batch size per device = {args.train_batch_size}")
1179
+ logger.info(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}")
1180
+ logger.info(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}")
1181
+ logger.info(f" Total optimization steps = {args.max_train_steps}")
1182
+ global_step = 0
1183
+ first_epoch = 0
1184
+
1185
+ # Potentially load in the weights and states from a previous save
1186
+ if args.resume_from_checkpoint:
1187
+ if args.resume_from_checkpoint != "latest":
1188
+ path = os.path.basename(args.resume_from_checkpoint)
1189
+ else:
1190
+ # Get the most recent checkpoint
1191
+ dirs = os.listdir(args.output_dir)
1192
+ dirs = [d for d in dirs if d.startswith("checkpoint")]
1193
+ dirs = sorted(dirs, key=lambda x: int(x.split("-")[1]))
1194
+ path = dirs[-1] if len(dirs) > 0 else None
1195
+
1196
+ if path is None:
1197
+ accelerator.print(
1198
+ f"Checkpoint '{args.resume_from_checkpoint}' does not exist. Starting a new training run."
1199
+ )
1200
+ args.resume_from_checkpoint = None
1201
+ initial_global_step = 0
1202
+ else:
1203
+ accelerator.print(f"Resuming from checkpoint {path}")
1204
+ accelerator.load_state(os.path.join(args.output_dir, path))
1205
+ global_step = int(path.split("-")[1])
1206
+
1207
+ initial_global_step = global_step
1208
+ first_epoch = global_step // num_update_steps_per_epoch
1209
+ else:
1210
+ initial_global_step = 0
1211
+
1212
+ progress_bar = tqdm(
1213
+ range(0, args.max_train_steps),
1214
+ initial=initial_global_step,
1215
+ desc="Steps",
1216
+ # Only show the progress bar once on each machine.
1217
+ disable=not accelerator.is_local_main_process,
1218
+ )
1219
+
1220
+ for epoch in range(first_epoch, args.num_train_epochs):
1221
+ unet.train()
1222
+ if args.train_text_encoder:
1223
+ text_encoder.train()
1224
+ for step, batch in enumerate(train_dataloader):
1225
+ with accelerator.accumulate(unet):
1226
+ pixel_values = batch["pixel_values"].to(dtype=weight_dtype)
1227
+
1228
+ if vae is not None:
1229
+ # Convert images to latent space
1230
+ model_input = vae.encode(batch["pixel_values"].to(dtype=weight_dtype)).latent_dist.sample()
1231
+ model_input = model_input * vae.config.scaling_factor
1232
+ else:
1233
+ model_input = pixel_values
1234
+
1235
+ # Sample noise that we'll add to the model input
1236
+ if args.offset_noise:
1237
+ noise = torch.randn_like(model_input) + 0.1 * torch.randn(
1238
+ model_input.shape[0], model_input.shape[1], 1, 1, device=model_input.device
1239
+ )
1240
+ else:
1241
+ noise = torch.randn_like(model_input)
1242
+ bsz, channels, height, width = model_input.shape
1243
+ # Sample a random timestep for each image
1244
+ timesteps = torch.randint(
1245
+ 0, noise_scheduler.config.num_train_timesteps, (bsz,), device=model_input.device
1246
+ )
1247
+ timesteps = timesteps.long()
1248
+
1249
+ # Add noise to the model input according to the noise magnitude at each timestep
1250
+ # (this is the forward diffusion process)
1251
+ noisy_model_input = noise_scheduler.add_noise(model_input, noise, timesteps)
1252
+
1253
+ # Get the text embedding for conditioning
1254
+ if args.pre_compute_text_embeddings:
1255
+ encoder_hidden_states = batch["input_ids"]
1256
+ else:
1257
+ encoder_hidden_states = encode_prompt(
1258
+ text_encoder,
1259
+ batch["input_ids"],
1260
+ batch["attention_mask"],
1261
+ text_encoder_use_attention_mask=args.text_encoder_use_attention_mask,
1262
+ )
1263
+
1264
+ if unwrap_model(unet).config.in_channels == channels * 2:
1265
+ noisy_model_input = torch.cat([noisy_model_input, noisy_model_input], dim=1)
1266
+
1267
+ if args.class_labels_conditioning == "timesteps":
1268
+ class_labels = timesteps
1269
+ else:
1270
+ class_labels = None
1271
+
1272
+ # Predict the noise residual
1273
+ model_pred = unet(
1274
+ noisy_model_input, timesteps, encoder_hidden_states, class_labels=class_labels, return_dict=False
1275
+ )[0]
1276
+
1277
+ if model_pred.shape[1] == 6:
1278
+ model_pred, _ = torch.chunk(model_pred, 2, dim=1)
1279
+
1280
+ # Get the target for loss depending on the prediction type
1281
+ if noise_scheduler.config.prediction_type == "epsilon":
1282
+ target = noise
1283
+ elif noise_scheduler.config.prediction_type == "v_prediction":
1284
+ target = noise_scheduler.get_velocity(model_input, noise, timesteps)
1285
+ else:
1286
+ raise ValueError(f"Unknown prediction type {noise_scheduler.config.prediction_type}")
1287
+
1288
+ if args.with_prior_preservation:
1289
+ # Chunk the noise and model_pred into two parts and compute the loss on each part separately.
1290
+ model_pred, model_pred_prior = torch.chunk(model_pred, 2, dim=0)
1291
+ target, target_prior = torch.chunk(target, 2, dim=0)
1292
+ # Compute prior loss
1293
+ prior_loss = F.mse_loss(model_pred_prior.float(), target_prior.float(), reduction="mean")
1294
+
1295
+ # Compute instance loss
1296
+ if args.snr_gamma is None:
1297
+ loss = F.mse_loss(model_pred.float(), target.float(), reduction="mean")
1298
+ else:
1299
+ # Compute loss-weights as per Section 3.4 of https://arxiv.org/abs/2303.09556.
1300
+ # Since we predict the noise instead of x_0, the original formulation is slightly changed.
1301
+ # This is discussed in Section 4.2 of the same paper.
1302
+ snr = compute_snr(noise_scheduler, timesteps)
1303
+
1304
+ if noise_scheduler.config.prediction_type == "v_prediction":
1305
+ # Velocity objective needs to be floored to an SNR weight of one.
1306
+ divisor = snr + 1
1307
+ else:
1308
+ divisor = snr
1309
+
1310
+ mse_loss_weights = (
1311
+ torch.stack([snr, args.snr_gamma * torch.ones_like(timesteps)], dim=1).min(dim=1)[0] / divisor
1312
+ )
1313
+
1314
+ loss = F.mse_loss(model_pred.float(), target.float(), reduction="none")
1315
+ loss = loss.mean(dim=list(range(1, len(loss.shape)))) * mse_loss_weights
1316
+ loss = loss.mean()
1317
+
1318
+ if args.with_prior_preservation:
1319
+ # Add the prior loss to the instance loss.
1320
+ loss = loss + args.prior_loss_weight * prior_loss
1321
+
1322
+ accelerator.backward(loss)
1323
+ if accelerator.sync_gradients:
1324
+ params_to_clip = (
1325
+ itertools.chain(unet.parameters(), text_encoder.parameters())
1326
+ if args.train_text_encoder
1327
+ else unet.parameters()
1328
+ )
1329
+ accelerator.clip_grad_norm_(params_to_clip, args.max_grad_norm)
1330
+ optimizer.step()
1331
+ lr_scheduler.step()
1332
+ optimizer.zero_grad(set_to_none=args.set_grads_to_none)
1333
+
1334
+ # Checks if the accelerator has performed an optimization step behind the scenes
1335
+ if accelerator.sync_gradients:
1336
+ progress_bar.update(1)
1337
+ global_step += 1
1338
+
1339
+ if accelerator.is_main_process:
1340
+ if global_step % args.checkpointing_steps == 0:
1341
+ # _before_ saving state, check if this save would set us over the `checkpoints_total_limit`
1342
+ if args.checkpoints_total_limit is not None:
1343
+ checkpoints = os.listdir(args.output_dir)
1344
+ checkpoints = [d for d in checkpoints if d.startswith("checkpoint")]
1345
+ checkpoints = sorted(checkpoints, key=lambda x: int(x.split("-")[1]))
1346
+
1347
+ # before we save the new checkpoint, we need to have at _most_ `checkpoints_total_limit - 1` checkpoints
1348
+ if len(checkpoints) >= args.checkpoints_total_limit:
1349
+ num_to_remove = len(checkpoints) - args.checkpoints_total_limit + 1
1350
+ removing_checkpoints = checkpoints[0:num_to_remove]
1351
+
1352
+ logger.info(
1353
+ f"{len(checkpoints)} checkpoints already exist, removing {len(removing_checkpoints)} checkpoints"
1354
+ )
1355
+ logger.info(f"removing checkpoints: {', '.join(removing_checkpoints)}")
1356
+
1357
+ for removing_checkpoint in removing_checkpoints:
1358
+ removing_checkpoint = os.path.join(args.output_dir, removing_checkpoint)
1359
+ shutil.rmtree(removing_checkpoint)
1360
+
1361
+ save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
1362
+ accelerator.save_state(save_path)
1363
+ logger.info(f"Saved state to {save_path}")
1364
+
1365
+ images = []
1366
+
1367
+ if args.validation_prompt is not None and global_step % args.validation_steps == 0:
1368
+ images = log_validation(
1369
+ unwrap_model(text_encoder) if text_encoder is not None else text_encoder,
1370
+ tokenizer,
1371
+ unwrap_model(unet),
1372
+ vae,
1373
+ args,
1374
+ accelerator,
1375
+ weight_dtype,
1376
+ global_step,
1377
+ validation_prompt_encoder_hidden_states,
1378
+ validation_prompt_negative_prompt_embeds,
1379
+ )
1380
+
1381
+ logs = {"loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]}
1382
+ progress_bar.set_postfix(**logs)
1383
+ accelerator.log(logs, step=global_step)
1384
+
1385
+ if global_step >= args.max_train_steps:
1386
+ break
1387
+
1388
+ # Create the pipeline using the trained modules and save it.
1389
+ accelerator.wait_for_everyone()
1390
+ if accelerator.is_main_process:
1391
+ pipeline_args = {}
1392
+
1393
+ if text_encoder is not None:
1394
+ pipeline_args["text_encoder"] = unwrap_model(text_encoder)
1395
+
1396
+ if args.skip_save_text_encoder:
1397
+ pipeline_args["text_encoder"] = None
1398
+
1399
+ pipeline = DiffusionPipeline.from_pretrained(
1400
+ args.pretrained_model_name_or_path,
1401
+ unet=unwrap_model(unet),
1402
+ revision=args.revision,
1403
+ variant=args.variant,
1404
+ **pipeline_args,
1405
+ )
1406
+
1407
+ # We train on the simplified learning objective. If we were previously predicting a variance, we need the scheduler to ignore it
1408
+ scheduler_args = {}
1409
+
1410
+ if "variance_type" in pipeline.scheduler.config:
1411
+ variance_type = pipeline.scheduler.config.variance_type
1412
+
1413
+ if variance_type in ["learned", "learned_range"]:
1414
+ variance_type = "fixed_small"
1415
+
1416
+ scheduler_args["variance_type"] = variance_type
1417
+
1418
+ pipeline.scheduler = pipeline.scheduler.from_config(pipeline.scheduler.config, **scheduler_args)
1419
+
1420
+ pipeline.save_pretrained(args.output_dir)
1421
+
1422
+ if args.push_to_hub:
1423
+ save_model_card(
1424
+ repo_id,
1425
+ images=images,
1426
+ base_model=args.pretrained_model_name_or_path,
1427
+ train_text_encoder=args.train_text_encoder,
1428
+ prompt=args.instance_prompt,
1429
+ repo_folder=args.output_dir,
1430
+ pipeline=pipeline,
1431
+ )
1432
+ upload_folder(
1433
+ repo_id=repo_id,
1434
+ folder_path=args.output_dir,
1435
+ commit_message="End of training",
1436
+ ignore_patterns=["step_*", "epoch_*"],
1437
+ )
1438
+
1439
+ accelerator.end_training()
1440
+
1441
+
1442
+ if __name__ == "__main__":
1443
+ args = parse_args()
1444
+ main(args)