(args)
| 4 | |
| 5 | |
| 6 | def sample_diffusion(args): |
| 7 | from utils.triplane_util import load_triplane_data, decompose_featmaps, save_triplane_data |
| 8 | from diffusion.script_util import create_model_and_diffusion_from_args |
| 9 | |
| 10 | # dist_util.setup_dist(args.gpu_id) |
| 11 | |
| 12 | src_data, sizes = load_triplane_data(encoding_feat_path(args.tag), device=dist_util.dev()) |
| 13 | |
| 14 | model, diffusion = create_model_and_diffusion_from_args(args) |
| 15 | model_path = diffusion_model_path(args.tag, args.ema_rate, args.diff_n_iters) |
| 16 | model.load_state_dict(dist_util.load_state_dict(model_path, map_location="cpu")) |
| 17 | model.to(dist_util.dev()).eval() |
| 18 | |
| 19 | sample_fn = ( |
| 20 | diffusion.p_sample_loop if not args.use_ddim else diffusion.ddim_sample_loop |
| 21 | ) |
| 22 | |
| 23 | result_dir = os.path.join(args.tag, args.output) |
| 24 | os.makedirs(result_dir, exist_ok=True) |
| 25 | |
| 26 | C = src_data.shape[0] |
| 27 | H, W, D = sizes |
| 28 | batch_size = args.diff_batch_size |
| 29 | H, W, D = int(H * args.resize[0]), int(W * args.resize[1]), int(D * args.resize[2]) |
| 30 | print("H, W, D:", H, W, D) |
| 31 | |
| 32 | result_paths = [] |
| 33 | for i in range(0, args.n_samples, batch_size): |
| 34 | bs = min(batch_size, args.n_samples - i) |
| 35 | out_shape = [bs, C, H + D, W + D] |
| 36 | |
| 37 | cond = {'H': H, 'W': W, 'D': D} |
| 38 | samples = sample_fn(model, out_shape, progress=True, model_kwargs=cond) |
| 39 | samples_xy, samples_xz, samples_yz = decompose_featmaps(samples, (H, W, D)) |
| 40 | samples_xy = samples_xy.detach().cpu().numpy() |
| 41 | samples_xz = samples_xz.detach().cpu().numpy() |
| 42 | samples_yz = samples_yz.detach().cpu().numpy() |
| 43 | |
| 44 | for j in range(bs): |
| 45 | save_path = os.path.join(result_dir, f"{i+j:03d}", "feat.npz") |
| 46 | save_triplane_data(save_path, samples_xy[j], samples_xz[j], samples_yz[j]) |
| 47 | result_paths.append(save_path) |
| 48 | return result_paths |
| 49 | |
| 50 | |
| 51 | def decode(args, paths): |
no test coverage detected