Create model and diffusion from args.
(args)
| 5 | |
| 6 | |
| 7 | def create_model_and_diffusion_from_args(args): |
| 8 | """ |
| 9 | Create model and diffusion from args. |
| 10 | """ |
| 11 | diffusion = create_gaussian_diffusion(**args_to_dict(args, diffusion_defaults().keys())) |
| 12 | |
| 13 | if type(args.channel_mult) is str: |
| 14 | args.channel_mult = tuple(int(ch_mult) for ch_mult in args.channel_mult.split(",")) |
| 15 | if args.diff_net_type == "unet_small": |
| 16 | model = TriplaneUNetModelSmall(**args_to_dict(args, diffusion_model_defaults().keys())) |
| 17 | elif args.diff_net_type == "unet_raw": |
| 18 | model = TriplaneUNetModelSmallRaw(**args_to_dict(args, diffusion_model_defaults().keys())) |
| 19 | return model, diffusion |
| 20 | |
| 21 | |
| 22 | def create_gaussian_diffusion( |
no test coverage detected