From bbc9a84f484bcd024a363d02768b09e3ce3886c4 Mon Sep 17 00:00:00 2001 From: arfafax Date: Fri, 3 Jul 2020 13:49:42 -0700 Subject: [PATCH 1/2] Adding batching to run_generator --- run_generator.py | 34 +++++++++++++++++++++++++--------- 1 file changed, 25 insertions(+), 9 deletions(-) diff --git a/run_generator.py b/run_generator.py index 2f154143..a94b0e69 100755 --- a/run_generator.py +++ b/run_generator.py @@ -11,12 +11,15 @@ import dnnlib.tflib as tflib import re import sys +import time +from tqdm import tqdm import pretrained_networks #---------------------------------------------------------------------------- -def generate_images(network_pkl, seeds, truncation_psi): +def generate_images(network_pkl, seeds, truncation_psi, minibatch_size=8): + mb = minibatch_size print('Loading networks from "%s"...' % network_pkl) _G, _D, Gs = pretrained_networks.load_networks(network_pkl) noise_vars = [var for name, var in Gs.components.synthesis.vars.items() if name.startswith('noise')] @@ -24,16 +27,28 @@ def generate_images(network_pkl, seeds, truncation_psi): Gs_kwargs = dnnlib.EasyDict() Gs_kwargs.output_transform = dict(func=tflib.convert_images_to_uint8, nchw_to_nhwc=True) Gs_kwargs.randomize_noise = False + Gs_kwargs.minibatch_size = mb if truncation_psi is not None: Gs_kwargs.truncation_psi = truncation_psi - - for seed_idx, seed in enumerate(seeds): - print('Generating image for seed %d (%d/%d) ...' % (seed, seed_idx, len(seeds))) - rnd = np.random.RandomState(seed) - z = rnd.randn(1, *Gs.input_shape[1:]) # [minibatch, component] + rnd = np.random.RandomState(seeds[0]) + z = rnd.randn(len(seeds), *Gs.input_shape[1:]) # [minibatch, component] + + tstart = time.time() + for i in tqdm(range(len(seeds)//mb + 1)): + zs = z[i*mb:min((i+1)*mb, len(seeds))] + print("\n") + print('Generating image for seeds %d to %d (%d/%d) ...' % (seeds[0]+i*mb, min(seeds[0]+(i+1) * mb, len(seeds)), i, len(seeds)//mb + 1)) + t0 = time.time() tflib.set_vars({var: rnd.randn(*var.shape.as_list()) for var in noise_vars}) # [height, width] - images = Gs.run(z, None, **Gs_kwargs) # [minibatch, height, width, channel] - PIL.Image.fromarray(images[0], 'RGB').save(dnnlib.make_run_dir_path('seed%04d.png' % seed)) + images = Gs.run(zs, None, **Gs_kwargs) # [minibatch, height, width, channel] + for j in range(len(images)): + seed = seeds[0]+i*mb+j + PIL.Image.fromarray(images[j], 'RGB').save(dnnlib.make_run_dir_path('seed%05d.png' % seed)) + elapsed = time.time() - t0 + print('{:.2f} images per second (in {:.3f}sec)'.format(len(images) / elapsed, elapsed)) + + elapsed = time.time() - tstart + print('{:.2f} images per second (in {:.3f}sec)'.format(len(seeds) / elapsed, elapsed)) #---------------------------------------------------------------------------- @@ -93,7 +108,7 @@ def _parse_num_range(s): range_re = re.compile(r'^(\d+)-(\d+)$') m = range_re.match(s) if m: - return range(int(m.group(1)), int(m.group(2))+1) + return list(range(int(m.group(1)), int(m.group(2))+1)) vals = s.split(',') return [int(x) for x in vals] @@ -131,6 +146,7 @@ def main(): parser_generate_images.add_argument('--network', help='Network pickle filename', dest='network_pkl', required=True) parser_generate_images.add_argument('--seeds', type=_parse_num_range, help='List of random seeds', required=True) parser_generate_images.add_argument('--truncation-psi', type=float, help='Truncation psi (default: %(default)s)', default=0.5) + parser_generate_images.add_argument('--minibatch_size', type=int, help='Minibatch size (default: %(default)s)', default=8) parser_generate_images.add_argument('--result-dir', help='Root directory for run results (default: %(default)s)', default='results', metavar='DIR') parser_style_mixing_example = subparsers.add_parser('style-mixing-example', help='Generate style mixing video') From dae710fef38b58fab49637ea526b24373910a353 Mon Sep 17 00:00:00 2001 From: arfafax Date: Fri, 3 Jul 2020 14:07:30 -0700 Subject: [PATCH 2/2] Fixing seed range. --- run_generator.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/run_generator.py b/run_generator.py index a94b0e69..31dc0062 100755 --- a/run_generator.py +++ b/run_generator.py @@ -37,7 +37,7 @@ def generate_images(network_pkl, seeds, truncation_psi, minibatch_size=8): for i in tqdm(range(len(seeds)//mb + 1)): zs = z[i*mb:min((i+1)*mb, len(seeds))] print("\n") - print('Generating image for seeds %d to %d (%d/%d) ...' % (seeds[0]+i*mb, min(seeds[0]+(i+1) * mb, len(seeds)), i, len(seeds)//mb + 1)) + print('Generating image for seeds %d to %d (%d/%d) ...' % (seeds[0]+i*mb, min(seeds[0]+(i+1) * mb, seeds[0] + len(seeds)), i, len(seeds)//mb + 1)) t0 = time.time() tflib.set_vars({var: rnd.randn(*var.shape.as_list()) for var in noise_vars}) # [height, width] images = Gs.run(zs, None, **Gs_kwargs) # [minibatch, height, width, channel]