aikenml/data_mining
0
1from __future__ import division2from __future__ import print_function3 4import argparse5import time6 7import torch8from spatial_correlation_sampler import SpatialCorrelationSampler9from tqdm import trange10 11TIME_SCALES = {'s': 1, 'ms': 1000, 'us': 1000000}12 13parser = argparse.ArgumentParser()14parser.add_argument('backend', choices=['cpu', 'cuda'], default='cuda')15parser.add_argument('-b', '--batch-size', type=int, default=16)16parser.add_argument('-k', '--kernel-size', type=int, default=3)17parser.add_argument('--patch', type=int, default=3)18parser.add_argument('--patch_dilation', type=int, default=2)19parser.add_argument('-c', '--channel', type=int, default=64)20parser.add_argument('--height', type=int, default=100)21parser.add_argument('-w', '--width', type=int, default=100)22parser.add_argument('-s', '--stride', type=int, default=2)23parser.add_argument('-p', '--pad', type=int, default=1)24parser.add_argument('--scale', choices=['s', 'ms', 'us'], default='us')25parser.add_argument('-r', '--runs', type=int, default=100)26parser.add_argument('--dilation', type=int, default=2)27parser.add_argument('-d', '--dtype', choices=['half', 'float', 'double'])28 29args = parser.parse_args()30 31device = torch.device(args.backend)32 33if args.dtype == 'half':34 dtype = torch.float1635elif args.dtype == 'float':36 dtype = torch.float3237else:38 dtype = torch.float6439 40 41input1 = torch.randn(args.batch_size,42 args.channel,43 args.height,44 args.width,45 dtype=dtype,46 device=device,47 requires_grad=True)48input2 = torch.randn_like(input1)49 50correlation_sampler = SpatialCorrelationSampler(51 args.kernel_size,52 args.patch,53 args.stride,54 args.pad,55 args.dilation,56 args.patch_dilation)57 58# Force CUDA initialization59output = correlation_sampler(input1, input2)60print(output.size())61output.mean().backward()62forward_min = float('inf')63forward_time = 064backward_min = float('inf')65backward_time = 066for _ in trange(args.runs):67 correlation_sampler.zero_grad()68 69 start = time.time()70 output = correlation_sampler(input1, input2)71 elapsed = time.time() - start72 forward_min = min(forward_min, elapsed)73 forward_time += elapsed74 output = output.mean()75 76 start = time.time()77 (output.mean()).backward()78 elapsed = time.time() - start79 backward_min = min(backward_min, elapsed)80 backward_time += elapsed81 82scale = TIME_SCALES[args.scale]83forward_min *= scale84backward_min *= scale85forward_average = forward_time / args.runs * scale86backward_average = backward_time / args.runs * scale87 88print('Forward: {0:.3f}/{1:.3f} {4} | Backward {2:.3f}/{3:.3f} {4}'.format(89 forward_min, forward_average, backward_min, backward_average,90 args.scale))91 