MCPcopy Create free account
hub / github.com/CSAILVision/gandissect / main

Function main

netdissect/evalablate.py:37–176  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

35'''
36
37def main():
38 # Training settings
39 def strpair(arg):
40 p = tuple(arg.split(':'))
41 if len(p) == 1:
42 p = p + p
43 return p
44
45 parser = argparse.ArgumentParser(description='Ablation eval',
46 epilog=textwrap.dedent(help_epilog),
47 formatter_class=argparse.RawDescriptionHelpFormatter)
48 parser.add_argument('--model', type=str, default=None,
49 help='constructor for the model to test')
50 parser.add_argument('--pthfile', type=str, default=None,
51 help='filename of .pth file for the model')
52 parser.add_argument('--outdir', type=str, default='dissect', required=True,
53 help='directory for dissection output')
54 parser.add_argument('--layers', type=strpair, nargs='+',
55 help='space-separated list of layer names to edit' +
56 ', in the form layername[:reportedname]')
57 parser.add_argument('--classes', type=str, nargs='+',
58 help='space-separated list of class names to ablate')
59 parser.add_argument('--metric', type=str, default='iou',
60 help='ordering metric for selecting units')
61 parser.add_argument('--unitcount', type=int, default=30,
62 help='number of units to ablate')
63 parser.add_argument('--segmenter', type=str,
64 help='directory containing segmentation dataset')
65 parser.add_argument('--netname', type=str, default=None,
66 help='name for network in generated reports')
67 parser.add_argument('--batch_size', type=int, default=5,
68 help='batch size for forward pass')
69 parser.add_argument('--size', type=int, default=200,
70 help='number of images to test')
71 parser.add_argument('--no-cuda', action='store_true', default=False,
72 help='disables CUDA usage')
73 parser.add_argument('--quiet', action='store_true', default=False,
74 help='silences console output')
75 if len(sys.argv) == 1:
76 parser.print_usage(sys.stderr)
77 sys.exit(1)
78 args = parser.parse_args()
79
80 # Set up console output
81 verbose_progress(not args.quiet)
82
83 # Speed up pytorch
84 torch.backends.cudnn.benchmark = True
85
86 # Set up CUDA
87 args.cuda = not args.no_cuda and torch.cuda.is_available()
88 if args.cuda:
89 torch.backends.cudnn.benchmark = True
90
91 # Take defaults for model constructor etc from dissect.json settings.
92 with open(os.path.join(args.outdir, 'dissect.json')) as f:
93 dissection = EasyDict(json.load(f))
94 if args.model is None:

Callers 1

evalablate.pyFile · 0.70

Calls 11

verbose_progressFunction · 0.90
EasyDictClass · 0.90
standard_z_sampleFunction · 0.90
autoimport_evalFunction · 0.90
default_progressFunction · 0.90
post_progressFunction · 0.90
measure_ablationFunction · 0.85
joinMethod · 0.80
parametersMethod · 0.45

Tested by

no test coverage detected