-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathvisualize.py
More file actions
68 lines (58 loc) · 2.67 KB
/
Copy pathvisualize.py
File metadata and controls
68 lines (58 loc) · 2.67 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
"""
D2NN unified visualization entrypoint.
"""
import argparse
from pathlib import Path
from tasks import run_classification_visualization, run_imaging_visualization
def build_parser():
parser = argparse.ArgumentParser(description="D2NN visualization")
parser.add_argument("--task", type=str, default="classification", choices=["classification", "imaging"])
parser.add_argument(
"--dataset",
type=str,
default="mnist",
help="classification: mnist/fashion-mnist/cifar10-gray/cifar10-rgb; imaging: stl10/imagefolder",
)
parser.add_argument("--checkpoint", type=str, required=True)
parser.add_argument("--image-root", type=str, default=None)
parser.add_argument("--size", type=int, default=None)
parser.add_argument("--layers", type=int, default=None)
parser.add_argument("--image-size", type=int, default=64)
parser.add_argument("--num-samples", type=int, default=6)
parser.add_argument("--input-fraction", type=float, default=0.5)
parser.add_argument("--seed", type=int, default=None, help="dataset split seed for imagefolder visualization")
parser.add_argument("--output-dir", type=str, default=None)
parser.add_argument("--no-show", action="store_true", help="Save figures without opening windows")
parser.add_argument("--understanding-report", action="store_true", help="also save understanding-report figures")
parser.add_argument("--sample-indices", type=str, default="0,1,2")
parser.add_argument("--quantization-levels", type=str, default="8,16")
parser.add_argument(
"--rs-backend",
type=str,
default=None,
choices=["direct", "fft"],
help="optional override for the RS propagation backend; defaults to the checkpoint manifest",
)
parser.add_argument(
"--propagation-chunk-size",
type=int,
default=None,
help="optional override for direct-backend chunk size; defaults to the checkpoint manifest",
)
parser.add_argument("--wavelength", type=float, default=None)
parser.add_argument("--layer-distance", type=float, default=None)
parser.add_argument("--pixel-size", type=float, default=None)
parser.add_argument("--input-distance", type=float, default=None)
parser.add_argument("--output-distance", type=float, default=None)
return parser
def main(argv=None):
args = build_parser().parse_args(argv)
if args.task == "imaging" and args.dataset == "mnist":
args.dataset = "stl10"
args.repo_root = Path(__file__).parent
if args.task == "classification":
run_classification_visualization(args)
else:
run_imaging_visualization(args)
if __name__ == "__main__":
main()