forked from nv-tlabs/vipe
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathsetup.py
More file actions
78 lines (63 loc) · 2.51 KB
/
Copy pathsetup.py
File metadata and controls
78 lines (63 loc) · 2.51 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
69
70
71
72
73
74
75
76
77
78
import os
import shutil
from pathlib import Path
from setuptools import find_packages, setup
from setuptools.command.build_py import build_py as _build_py
try:
import torch
import torch.version
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
cuda_version = torch.version.cuda
assert cuda_version is not None, "Pytorch CUDA is required for this installation."
except ImportError:
raise ValueError("Pytorch not found, please install it first.")
PACKAGE_NAME = "vipe"
SOURCE_CONFIG_DIR = Path(__file__).resolve().parent / "configs"
coder_finder_path = f"{PACKAGE_NAME}/ext/specs.py"
code_finder_namespace = {"__file__": coder_finder_path}
with open(coder_finder_path, "r") as fh:
exec(fh.read(), code_finder_namespace)
get_sources = code_finder_namespace["get_sources"]
get_cpp_flags = code_finder_namespace["get_cpp_flags"]
get_cuda_flags = code_finder_namespace["get_cuda_flags"]
class build_py(_build_py):
def run(self) -> None:
super().run()
self._copy_configs()
def _copy_configs(self) -> None:
if not SOURCE_CONFIG_DIR.is_dir():
raise RuntimeError(f"Missing config source directory: {SOURCE_CONFIG_DIR}")
target = Path(self.build_lib) / PACKAGE_NAME / "_configs"
if target.exists():
shutil.rmtree(target)
shutil.copytree(
SOURCE_CONFIG_DIR,
target,
ignore=shutil.ignore_patterns("__pycache__", "*.pyc"),
)
(target / "__init__.py").write_text(
"# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.\n"
"# SPDX-License-Identifier: Apache-2.0\n\n"
'"""Package data target for build-time generated ViPE configs."""\n',
encoding="utf-8",
)
# Setup CUDA_HOME for conda environment for consistency
if "CONDA_PREFIX" in os.environ:
conda_nvcc_path = os.path.join(os.environ["CONDA_PREFIX"], "bin", "nvcc")
if os.path.exists(conda_nvcc_path):
os.environ["PYTORCH_NVCC"] = conda_nvcc_path
cpp_flags = get_cpp_flags()
cuda_flags = get_cuda_flags()
packages = find_packages()
setup(
packages=packages,
include_package_data=True,
ext_modules=[
CUDAExtension(
f"{PACKAGE_NAME}_ext",
sources=get_sources(), # type: ignore
extra_compile_args={"cxx": cpp_flags, "nvcc": cuda_flags}, # type: ignore
)
],
cmdclass={"build_ext": BuildExtension.with_options(use_ninja=True), "build_py": build_py},
)