MCPcopy Create free account
hub / github.com/ChenWu98/cycle-diffusion / prepare_latentdiff

Function prepare_latentdiff

model/gan_wrapper/latentdiff_wrapper.py:16–55  ·  view source on GitHub ↗
(source_model_type, custom_steps)

Source from the content-addressed store, hash-verified

14
15
16def prepare_latentdiff(source_model_type, custom_steps):
17 print('First of all, when the code changes, make sure that no part in the model is under no_grad!')
18 latentdiff_ckpt = os.path.join('ckpts', 'ldm_models', 'ldm', source_model_type, 'model.ckpt')
19
20 parser = get_parser()
21 opt, unknown = parser.parse_known_args(
22 [
23 '-r', latentdiff_ckpt,
24 '-c', str(custom_steps),
25 '-e', str(0),
26 ]
27 )
28 print('unknown args: ', unknown)
29
30 if not os.path.exists(opt.resume):
31 raise ValueError("Cannot find {}".format(opt.resume))
32 if os.path.isfile(opt.resume):
33 # paths = opt.resume.split("/")
34 try:
35 logdir = '/'.join(opt.resume.split('/')[:-1])
36 # idx = len(paths)-paths[::-1].index("logs")+1
37 print(f'Logdir is {logdir}')
38 except ValueError:
39 paths = opt.resume.split("/")
40 idx = -2 # take a guess: path/to/logdir/checkpoints/model.ckpt
41 logdir = "/".join(paths[:idx])
42 ckpt = opt.resume
43 else:
44 assert os.path.isdir(opt.resume), f"{opt.resume} is not a directory"
45 logdir = opt.resume.rstrip("/")
46 ckpt = os.path.join(logdir, "model.ckpt")
47
48 base_configs = sorted(glob.glob(os.path.join(logdir, "config.yaml")))
49 opt.base = base_configs
50
51 configs = [OmegaConf.load(cfg) for cfg in opt.base]
52 cli = OmegaConf.from_dotlist(unknown)
53 config = OmegaConf.merge(*configs, cli)
54
55 return config, opt, ckpt
56
57
58def get_condition(model, class_label, bs):

Callers 1

__init__Method · 0.70

Calls 1

get_parserFunction · 0.90

Tested by

no test coverage detected