MCPcopy Create free account
hub / github.com/YZY-stack/DF40 / main

Function main

EFS_finetune_code/diffusion_based/train_scripts/train_pixart_lora.py:422–970  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

420
421
422def 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

Callers 1

Calls 9

parseMethod · 0.80
encodeMethod · 0.80
backwardMethod · 0.80
stepMethod · 0.80
parse_argsFunction · 0.70
ImageTextDatasetClass · 0.70
save_model_cardFunction · 0.70
trainMethod · 0.45
updateMethod · 0.45

Tested by

no test coverage detected