Team Ai
Apppublic

aikenml/data_mining

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
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