()
| 420 | |
| 421 | |
| 422 | def main(): |
| 423 | args = parse_args() |
| 424 | logging_dir = Path(args.output_dir, args.logging_dir) |
| 425 | |
| 426 | accelerator_project_config = ProjectConfiguration(project_dir=args.output_dir, logging_dir=logging_dir) |
| 427 | |
| 428 | accelerator = Accelerator( |
| 429 | gradient_accumulation_steps=args.gradient_accumulation_steps, |
| 430 | mixed_precision=args.mixed_precision, |
| 431 | log_with=args.report_to, |
| 432 | project_config=accelerator_project_config, |
| 433 | ) |
| 434 | if args.report_to == "wandb": |
| 435 | if not is_wandb_available(): |
| 436 | raise ImportError("Make sure to install wandb if you want to use it for logging during training.") |
| 437 | import wandb |
| 438 | |
| 439 | # Make one log on every process with the configuration for debugging. |
| 440 | logging.basicConfig( |
| 441 | format="%(asctime)s - %(levelname)s - %(name)s - %(message)s", |
| 442 | datefmt="%m/%d/%Y %H:%M:%S", |
| 443 | level=logging.INFO, |
| 444 | ) |
| 445 | logger.info(accelerator.state, main_process_only=False) |
| 446 | if accelerator.is_local_main_process: |
| 447 | datasets.utils.logging.set_verbosity_warning() |
| 448 | transformers.utils.logging.set_verbosity_warning() |
| 449 | diffusers.utils.logging.set_verbosity_info() |
| 450 | else: |
| 451 | datasets.utils.logging.set_verbosity_error() |
| 452 | transformers.utils.logging.set_verbosity_error() |
| 453 | diffusers.utils.logging.set_verbosity_error() |
| 454 | |
| 455 | # If passed along, set the training seed now. |
| 456 | if args.seed is not None: |
| 457 | set_seed(args.seed) |
| 458 | |
| 459 | # Handle the repository creation |
| 460 | if accelerator.is_main_process: |
| 461 | if args.output_dir is not None: |
| 462 | os.makedirs(args.output_dir, exist_ok=True) |
| 463 | |
| 464 | if args.push_to_hub: |
| 465 | repo_id = create_repo(repo_id=args.hub_model_id or Path(args.output_dir).name, exist_ok=True, token=args.hub_token).repo_id |
| 466 | |
| 467 | # See Section 3.1. of the paper. |
| 468 | max_length = 120 |
| 469 | |
| 470 | # Load scheduler, tokenizer and models. |
| 471 | noise_scheduler = DDPMScheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="scheduler") |
| 472 | tokenizer = T5Tokenizer.from_pretrained(args.pretrained_model_name_or_path, subfolder="tokenizer", revision=args.revision) |
| 473 | |
| 474 | text_encoder = T5EncoderModel.from_pretrained(args.pretrained_model_name_or_path, subfolder="text_encoder", revision=args.revision) |
| 475 | |
| 476 | vae = AutoencoderKL.from_pretrained(args.pretrained_model_name_or_path, subfolder="vae", revision=args.revision, variant=args.variant) |
| 477 | |
| 478 | transformer = Transformer2DModel.from_pretrained(args.pretrained_model_name_or_path, subfolder="transformer", torch_dtype=torch.float16) |
| 479 |
no test coverage detected