setup xpu ops
()
| 124 | |
| 125 | |
| 126 | def xpu_setup_ops(): |
| 127 | """ |
| 128 | setup xpu ops |
| 129 | """ |
| 130 | PADDLE_PATH = os.path.dirname(paddle.__file__) |
| 131 | PADDLE_INCLUDE_PATH = os.path.join(PADDLE_PATH, "include") |
| 132 | PADDLE_LIB_PATH = os.path.join(PADDLE_PATH, "libs") |
| 133 | |
| 134 | BKCL_PATH = os.getenv("BKCL_PATH") |
| 135 | if BKCL_PATH is None: |
| 136 | BKCL_INC_PATH = os.path.join(PADDLE_INCLUDE_PATH, "xpu") |
| 137 | BKCL_LIB_PATH = os.path.join(PADDLE_LIB_PATH, "libbkcl.so") |
| 138 | else: |
| 139 | BKCL_INC_PATH = os.path.join(BKCL_PATH, "include") |
| 140 | BKCL_LIB_PATH = os.path.join(BKCL_PATH, "so", "libbkcl.so") |
| 141 | |
| 142 | CLANG_PATH = os.getenv("CLANG_PATH") |
| 143 | assert CLANG_PATH is not None, "CLANG_PATH is not set." |
| 144 | |
| 145 | XRE_PATH = os.getenv("XRE_PATH") |
| 146 | if XRE_PATH is None: |
| 147 | XRE_INC_PATH = os.path.join(PADDLE_INCLUDE_PATH, "xre") |
| 148 | XRE_LIB_PATH = os.path.join(PADDLE_LIB_PATH, "libxpucuda.so") |
| 149 | XRE_LIB_DIR = os.path.join(PADDLE_LIB_PATH) |
| 150 | else: |
| 151 | XRE_INC_PATH = os.path.join(XRE_PATH, "include") |
| 152 | XRE_LIB_PATH = os.path.join(XRE_PATH, "so", "libxpucuda.so") |
| 153 | XRE_LIB_DIR = os.path.join(XRE_PATH, "so") |
| 154 | |
| 155 | XDNN_PATH = os.getenv("XDNN_PATH") |
| 156 | if XDNN_PATH is None: |
| 157 | XHPC_VERSION = paddle.version.xpu_xhpc() |
| 158 | print(f"Fetched XHPC_VERSION from paddle: {XHPC_VERSION}") |
| 159 | |
| 160 | XHPC_URL = f"https://klx-sdk-release-public.su.bcebos.com/xhpc/{XHPC_VERSION}/xhpc-ubuntu2004_x86_64.tar.gz" |
| 161 | THIRD_PARTY_PATH = os.path.join(current_file.parent, "third_party") |
| 162 | XHPC_PATH = os.path.join(THIRD_PARTY_PATH, "xhpc-ubuntu2004_x86_64") |
| 163 | |
| 164 | if os.path.exists(XHPC_PATH): |
| 165 | with open(os.path.join(XHPC_PATH, "version.txt")) as f: |
| 166 | date_line = [line.strip() for line in f.readlines() if "Date:" in line][0] |
| 167 | LOCAL_VERSION = f"dev/{date_line.split()[1]}" |
| 168 | if LOCAL_VERSION == XHPC_VERSION: |
| 169 | print("Local XHPC exists, skip downloading it again.") |
| 170 | else: |
| 171 | XHPC_UPDATE_POLICY_ENV = os.getenv("XHPC_UPDATE_POLICY") |
| 172 | if XHPC_UPDATE_POLICY_ENV is not None: |
| 173 | if XHPC_UPDATE_POLICY_ENV == "FORCE": |
| 174 | print("Forced update detected, downloading new XHPC.") |
| 175 | download_and_extract(XHPC_URL, THIRD_PARTY_PATH) |
| 176 | elif XHPC_UPDATE_POLICY_ENV == "SKIP": |
| 177 | print("Skipped updating XHPC.") |
| 178 | else: |
| 179 | raise Exception( |
| 180 | f"\033[91mInvalid value for environment variable XHPC_UPDATE_POLICY\033[0m: {XHPC_UPDATE_POLICY_ENV}, " |
| 181 | f"Valid environment values are FORCE or SKIP.", |
| 182 | ) |
| 183 | else: |
no test coverage detected