MCPcopy Create free account
hub / github.com/DPS2022/diffusion-posterior-sampling / main

Function main

sample_condition.py:25–118  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

23
24
25def 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)

Callers 1

Calls 12

get_loggerFunction · 0.90
create_modelFunction · 0.90
get_operatorFunction · 0.90
get_noiseFunction · 0.90
get_conditioning_methodFunction · 0.90
create_samplerFunction · 0.90
get_datasetFunction · 0.90
get_dataloaderFunction · 0.90
mask_generatorClass · 0.90
clear_colorFunction · 0.90
load_yamlFunction · 0.85
forwardMethod · 0.45

Tested by

no test coverage detected