File size: 2,587 Bytes
6d35aff
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
"""Check the CUDA/APE runtime needed by Orienter Stage 2."""

import argparse
import importlib
import importlib.metadata
import json
import platform
import sys


HISTORICAL_REFERENCE = {
    "python": "3.9.18",
    "torch": "2.2.0+cu121",
    "torchvision": "0.17.0+cu121",
    "detectron2": "0.6",
    "detrex": "0.3.0",
}


def _version(module_name, distribution_name=None):
    module = importlib.import_module(module_name)
    value = getattr(module, "__version__", None)
    if value is not None:
        return str(value)
    return importlib.metadata.version(distribution_name or module_name)


def check_environment(require_cuda=True, require_ape_extension=True):
    report = {
        "python": platform.python_version(),
        "historical_reference": dict(HISTORICAL_REFERENCE),
        "versions": {},
        "cuda_available": False,
        "ape_extension": False,
        "errors": [],
    }

    for module_name, distribution_name in (
        ("torch", "torch"),
        ("torchvision", "torchvision"),
        ("detectron2", "detectron2"),
        ("detrex", "detrex"),
    ):
        try:
            report["versions"][module_name] = _version(module_name, distribution_name)
        except Exception as exc:
            report["errors"].append(f"cannot import {module_name}: {type(exc).__name__}: {exc}")

    try:
        torch = importlib.import_module("torch")
        report["torch_cuda_runtime"] = torch.version.cuda
        report["cuda_available"] = bool(torch.cuda.is_available())
        if require_cuda and not report["cuda_available"]:
            report["errors"].append("torch.cuda.is_available() is false")
    except Exception:
        pass

    if require_ape_extension:
        try:
            importlib.import_module("ape._C")
            report["ape_extension"] = True
        except Exception as exc:
            report["errors"].append(
                f"cannot import compiled ape._C extension: {type(exc).__name__}: {exc}"
            )

    return report


def build_parser():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--allow-no-cuda", action="store_true")
    parser.add_argument("--skip-ape-extension", action="store_true")
    return parser


def main(argv=None):
    args = build_parser().parse_args(argv)
    report = check_environment(
        require_cuda=not args.allow_no_cuda,
        require_ape_extension=not args.skip_ape_extension,
    )
    print(json.dumps(report, indent=2, sort_keys=True))
    return 1 if report["errors"] else 0


if __name__ == "__main__":
    sys.exit(main())