MCPcopy Create free account
hub / github.com/NVlabs/DiffusionNFT / get_config

Function get_config

config/base.py:4–108  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

2
3
4def 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 ######

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected