MCPcopy Create free account
hub / github.com/Zhiyuan-R/Tiger-Diffusion / parse_args

Function parse_args

train_generation.py:781–844  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

779
780
781def parse_args():
782
783 parser = argparse.ArgumentParser()
784 parser.add_argument('--dataroot', default='ShapeNetCore.v2.PC15k/')
785 parser.add_argument('--category', default='chair')
786
787 parser.add_argument('--bs', type=int, default=64, help='input batch size')
788 parser.add_argument('--workers', type=int, default=16, help='workers')
789 parser.add_argument('--niter', type=int, default=10000, help='number of epochs to train for')
790
791 parser.add_argument('--nc', default=3)
792 parser.add_argument('--npoints', default=2048)
793 '''model'''
794 parser.add_argument('--beta_start', default=0.0001)
795 parser.add_argument('--beta_end', default=0.02)
796 parser.add_argument('--schedule_type', default='linear')
797 parser.add_argument('--time_num', default=1000)
798
799 #params
800 parser.add_argument('--attention', default=True)
801 parser.add_argument('--dropout', default=0.1)
802 parser.add_argument('--embed_dim', type=int, default=64)
803 parser.add_argument('--loss_type', default='mse')
804 parser.add_argument('--model_mean_type', default='eps')
805 parser.add_argument('--model_var_type', default='fixedsmall')
806
807 parser.add_argument('--lr', type=float, default=2e-4, help='learning rate for E, default=0.0002')
808 parser.add_argument('--beta1', type=float, default=0.5, help='beta1 for adam. default=0.5')
809 parser.add_argument('--decay', type=float, default=0, help='weight decay for EBM')
810 parser.add_argument('--grad_clip', type=float, default=None, help='weight decay for EBM')
811 parser.add_argument('--lr_gamma', type=float, default=0.998, help='lr decay for EBM')
812
813 parser.add_argument('--model', default='', help="path to model (to continue training)")
814
815
816 '''distributed'''
817 parser.add_argument('--world_size', default=1, type=int,
818 help='Number of distributed nodes.')
819 parser.add_argument('--dist_url', default='tcp://127.0.0.1:9991', type=str,
820 help='url used to set up distributed training')
821 parser.add_argument('--dist_backend', default='nccl', type=str,
822 help='distributed backend')
823 parser.add_argument('--distribution_type', default='single', choices=['multi', 'single', None],
824 help='Use multi-processing distributed training to launch '
825 'N processes per node, which has N GPUs. This is the '
826 'fastest way to use PyTorch for either single node or '
827 'multi node data parallel training')
828 parser.add_argument('--rank', default=0, type=int,
829 help='node rank for distributed training')
830 parser.add_argument('--gpu', default=None, type=int,
831 help='GPU id to use. None means using all available GPUs.')
832
833 '''eval'''
834 parser.add_argument('--saveIter', default=100, help='unit: epoch')
835 parser.add_argument('--diagIter', default=50, help='unit: epoch')
836 parser.add_argument('--vizIter', default=50, help='unit: epoch')
837 parser.add_argument('--print_freq', default=50, help='unit: iter')
838

Callers 1

mainFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected