(input_args=None)
| 46 | import wandb |
| 47 | |
| 48 | def parse_args(input_args=None): |
| 49 | parser = argparse.ArgumentParser(description="Argparser for ImageDream (diffusers) training script.") |
| 50 | |
| 51 | parser.add_argument("--seed", type=int, default=42, help="A seed for reproducible training.") |
| 52 | parser.add_argument("--guidance_scale", type=float, default=5.0) |
| 53 | parser.add_argument("--conditioning_dropout_prob", type=float, default=0.1, |
| 54 | help="Conditioning dropout probability. Drops out the conditionings (image and edit prompt) used in training InstructPix2Pix. See section 3.2.1 in the paper: https://arxiv.org/abs/2211.09800" |
| 55 | ) |
| 56 | parser.add_argument("--checkpointing_steps", type=int, default=100) |
| 57 | parser.add_argument("--checkpoints_total_limit", type=int, default=10, help=("Max number of checkpoints to store.")) |
| 58 | parser.add_argument("--resume_from_checkpoint", type=str, default=None) |
| 59 | parser.add_argument("--gradient_checkpointing", action="store_true", help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.") |
| 60 | parser.add_argument("--use_8bit_adam", action="store_true", help="Whether or not to use 8-bit Adam from bitsandbytes.") |
| 61 | parser.add_argument("--dataloader_num_workers", type=int, default=1) |
| 62 | parser.add_argument("--adam_beta1", type=float, default=0.9, help="The beta1 parameter for the Adam optimizer.") |
| 63 | parser.add_argument("--adam_beta2", type=float, default=0.999, help="The beta2 parameter for the Adam optimizer.") |
| 64 | parser.add_argument("--adam_weight_decay", type=float, default=1e-2, help="Weight decay to use.") |
| 65 | parser.add_argument("--adam_epsilon", type=float, default=1e-08, help="Epsilon value for the Adam optimizer") |
| 66 | parser.add_argument("--max_grad_norm", default=0.5, type=float, help="Max gradient norm.") |
| 67 | parser.add_argument("--push_to_hub", action="store_true", help="Whether or not to push the model to the Hub.") |
| 68 | parser.add_argument("--hub_token", type=str, default=None, help="The token to use to push to the Model Hub.") |
| 69 | parser.add_argument("--logging_dir", type=str, default="logs") |
| 70 | parser.add_argument("--report_to", type=str, default="wandb") |
| 71 | parser.add_argument("--set_grads_to_none", default=True) |
| 72 | parser.add_argument("--use_ema", action="store_true", help="Whether to use EMA model.") |
| 73 | |
| 74 | if input_args is not None: |
| 75 | args = parser.parse_args(input_args) |
| 76 | else: |
| 77 | args = parser.parse_args() |
| 78 | |
| 79 | return args |
| 80 | |
| 81 | def _encode_text_prompt( |
| 82 | tokenizer, |
no outgoing calls
no test coverage detected