()
| 23 | |
| 24 | |
| 25 | def main(): |
| 26 | parser = argparse.ArgumentParser() |
| 27 | parser.add_argument('--model_config', type=str) |
| 28 | parser.add_argument('--diffusion_config', type=str) |
| 29 | parser.add_argument('--task_config', type=str) |
| 30 | parser.add_argument('--gpu', type=int, default=0) |
| 31 | parser.add_argument('--save_dir', type=str, default='./results') |
| 32 | args = parser.parse_args() |
| 33 | |
| 34 | # logger |
| 35 | logger = get_logger() |
| 36 | |
| 37 | # Device setting |
| 38 | device_str = f"cuda:{args.gpu}" if torch.cuda.is_available() else 'cpu' |
| 39 | logger.info(f"Device set to {device_str}.") |
| 40 | device = torch.device(device_str) |
| 41 | |
| 42 | # Load configurations |
| 43 | model_config = load_yaml(args.model_config) |
| 44 | diffusion_config = load_yaml(args.diffusion_config) |
| 45 | task_config = load_yaml(args.task_config) |
| 46 | |
| 47 | #assert model_config['learn_sigma'] == diffusion_config['learn_sigma'], \ |
| 48 | #"learn_sigma must be the same for model and diffusion configuartion." |
| 49 | |
| 50 | # Load model |
| 51 | model = create_model(**model_config) |
| 52 | model = model.to(device) |
| 53 | model.eval() |
| 54 | |
| 55 | # Prepare Operator and noise |
| 56 | measure_config = task_config['measurement'] |
| 57 | operator = get_operator(device=device, **measure_config['operator']) |
| 58 | noiser = get_noise(**measure_config['noise']) |
| 59 | logger.info(f"Operation: {measure_config['operator']['name']} / Noise: {measure_config['noise']['name']}") |
| 60 | |
| 61 | # Prepare conditioning method |
| 62 | cond_config = task_config['conditioning'] |
| 63 | cond_method = get_conditioning_method(cond_config['method'], operator, noiser, **cond_config['params']) |
| 64 | measurement_cond_fn = cond_method.conditioning |
| 65 | logger.info(f"Conditioning method : {task_config['conditioning']['method']}") |
| 66 | |
| 67 | # Load diffusion sampler |
| 68 | sampler = create_sampler(**diffusion_config) |
| 69 | sample_fn = partial(sampler.p_sample_loop, model=model, measurement_cond_fn=measurement_cond_fn) |
| 70 | |
| 71 | # Working directory |
| 72 | out_path = os.path.join(args.save_dir, measure_config['operator']['name']) |
| 73 | os.makedirs(out_path, exist_ok=True) |
| 74 | for img_dir in ['input', 'recon', 'progress', 'label']: |
| 75 | os.makedirs(os.path.join(out_path, img_dir), exist_ok=True) |
| 76 | |
| 77 | # Prepare dataloader |
| 78 | data_config = task_config['data'] |
| 79 | transform = transforms.Compose([transforms.ToTensor(), |
| 80 | transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))]) |
| 81 | dataset = get_dataset(**data_config, transforms=transform) |
| 82 | loader = get_dataloader(dataset, batch_size=1, num_workers=0, train=False) |
no test coverage detected