()
| 2 | |
| 3 | |
| 4 | def get_config(): |
| 5 | config = ml_collections.ConfigDict() |
| 6 | |
| 7 | ###### General ###### |
| 8 | # run name for wandb logging and checkpoint saving -- if not provided, will be auto-generated based on the datetime. |
| 9 | config.run_name = "" |
| 10 | config.debug = False |
| 11 | |
| 12 | # random seed for reproducibility. |
| 13 | config.seed = 42 |
| 14 | # top-level logging directory for checkpoint saving. |
| 15 | config.logdir = "logs" |
| 16 | # number of epochs to train for. each epoch is one round of sampling from the model followed by training on those |
| 17 | # samples. |
| 18 | config.num_epochs = 100000 |
| 19 | # number of epochs between saving model checkpoints. |
| 20 | config.save_freq = 30 |
| 21 | config.eval_freq = 10 |
| 22 | # mixed precision training. options are "fp16", "bf16", and "no". half-precision speeds up training significantly. |
| 23 | config.mixed_precision = "fp16" |
| 24 | # allow tf32 on Ampere GPUs, which can speed up training. |
| 25 | config.allow_tf32 = True |
| 26 | # resume training from a checkpoint. either an exact checkpoint directory (e.g. checkpoint_50), or a directory |
| 27 | # containing checkpoints, in which case the latest one will be used. `config.use_lora` must be set to the same value |
| 28 | # as the run that generated the saved checkpoint. |
| 29 | config.resume_from = "" |
| 30 | # whether or not to use LoRA. |
| 31 | config.use_lora = True |
| 32 | config.dataset = "" |
| 33 | config.resolution = 768 |
| 34 | |
| 35 | ###### Pretrained Model ###### |
| 36 | config.pretrained = pretrained = ml_collections.ConfigDict() |
| 37 | # base model to load. either a path to a local directory, or a model name from the HuggingFace model hub. |
| 38 | pretrained.model = "" |
| 39 | # revision of the model to load. |
| 40 | pretrained.revision = "" |
| 41 | |
| 42 | ###### Sampling ###### |
| 43 | config.sample = sample = ml_collections.ConfigDict() |
| 44 | # number of sampler inference steps. |
| 45 | sample.num_steps = 40 |
| 46 | sample.eval_num_steps = 40 |
| 47 | # classifier-free guidance weight. 1.0 is no guidance. |
| 48 | sample.guidance_scale = 4.5 |
| 49 | # batch size (per GPU!) to use for sampling. |
| 50 | sample.train_batch_size = 1 |
| 51 | sample.num_image_per_prompt = 1 |
| 52 | sample.test_batch_size = 1 |
| 53 | # number of batches to sample per epoch. the total number of samples per epoch is `num_batches_per_epoch * |
| 54 | # batch_size * num_gpus`. |
| 55 | sample.num_batches_per_epoch = 2 |
| 56 | # Whether use all samples in a batch to compute std |
| 57 | sample.global_std = True |
| 58 | # noise level |
| 59 | sample.noise_level = 1.0 |
| 60 | |
| 61 | ###### Training ###### |
nothing calls this directly
no outgoing calls
no test coverage detected