diff --git a/.claude/CLAUDE.md b/.claude/CLAUDE.md new file mode 100644 index 0000000..80f0f55 --- /dev/null +++ b/.claude/CLAUDE.md @@ -0,0 +1,111 @@ +# CLAUDE.md + +This file provides guidance to Claude Code (claude.ai/code) when working with code in this repository. + +## Project Overview + +VILA (Vision Infrastructure Library Assembler) is a C++ infrastructure library for computer vision applications. It provides JSON-based configuration, logging with WPP support, reflection registry, error handling (Status/StatusOr), profiling tools, and a template DAG library. + +## Build Systems + +### Bazel (Primary) + +```bash +# Build all tests +bazelisk test //tests/clim/... //tests/vila/... + +# Build a specific test +bazelisk test //tests/clim:argparse_test + +# Build with specific config (e.g., SYCL) +bazelisk build --config=sycl //... + +# Run using legacy WORKSPACE (on Windows) +bazelisk --output_base="C:/temp/_vila_workspace" build --noenable_bzlmod //... +``` + +**Legacy WORKSPACE setup (Windows)** requires loading three workspace files in sequence: +```bazel +load("@vila//vila:workspace0.bzl", vila_workspace0 = "workspace") +vila_workspace0() +load("@vila//vila:workspace1.bzl", vila_workspace1 = "workspace") +vila_workspace1() +load("@vila//vila:workspace2.bzl", vila_workspace2 = "workspace") +``` + +### CMake + +```bash +mkdir build && cmake -Bbuild -S. -GNinja +cmake --build build --config Release + +# With testing enabled +cmake -Bbuild -S. -GNinja -DVILA_ENABLE_TESTING=ON +``` + +## Code Style + +- **Style**: Based on Google C++ Style (see `.clang-format`) +- **Clang-tidy**: Configured in `.clang-tidy`, targets `clim/` and `vila/` headers only (excludes generated files) +- **Pre-commit hooks**: Run via `pre-commit run -s HEAD^ -o HEAD` (see `.pre-commit-config.yaml`) +- **Spell checking**: Codespell configured with words bag at `.github/WORDS_BAG.txt` + +## Architecture + +``` +clim/ # Header-only utility library (math, strings, containers, etc.) + ├── argparse/ # Command-line argument parsing + ├── container/ # Bounding boxes, ring buffers, etc. + ├── filter/ # Kalman and alpha-beta filters + ├── hash/ # CityHash, MurmurHash + ├── math/ # Quaternion, numerical utilities + ├── os/ # OS utilities (aligned malloc, barriers) + ├── path/ # Cross-platform path handling + ├── reflection/ # Reflection registry + ├── string/ # String splitting, stripping, const_string + ├── vt/ # Vector math (GEMM, neural network ops) + └── zip/ # Zip utility functions + +vila/ # Core library components + ├── config/ # JSON-based configuration system + ├── graph/ # Template header-only DAG (dag.h, graph.h, route.h, traversal.h) + ├── hook/ # Windows DLL hooking (detours) + ├── logging/ # Logger with WPP support (code_location, logger) + ├── profiling/ # ITT, timer, trace utilities + ├── status/ # Status and StatusOr error handling + ├── widget/ # (UI components) + └── bazel/ # Bazel-specific build rules and toolchains + +python/ # Python bindings via nanobind/pybind11 +tests/ # GoogleTest-based C++ tests +``` + +## Key Dependencies (via Bazel) + +- `fmt` (12.1.0) - Formatting library +- `spdlog` - Logging library +- `googletest` - Testing framework +- `google_benchmark` - Benchmarking +- `rules_foreign_cc` - CMake/ ninja build support + +## Testing + +```bash +# C++ tests (Bazel) +bazelisk test //tests/clim/... //tests/vila/... + +# Python tests +pip install -e python[test] +pytest --cov=python/vila python/tests +``` + +## Windows-Specific Notes + +- Default C++ standard: C++17 +- Windows WPP logging disabled by default (see commit ed44f79) +- Windows-specific configs use `select()` with `//conditions:default` since some build configs are Windows-only +- The `range_test` is Windows-only and requires C++20 + +## Editor Setup + +The project includes `.vscode/` settings for convenience with bazelized projects. diff --git a/.github/workflows/build-tests.yml b/.github/workflows/build-tests.yml index 12a2acf..9903569 100644 --- a/.github/workflows/build-tests.yml +++ b/.github/workflows/build-tests.yml @@ -35,6 +35,9 @@ jobs: uses: actions/setup-python@v5 with: python-version: "3.12" + - name: Setup system python env + run: | + python -m pip install -r tests/vila/requirements_tvm_ffi_test.txt - uses: bazel-contrib/setup-bazel@0.15.0 with: bazelisk-cache: true diff --git a/.gitignore b/.gitignore index 4463ce4..034d91b 100644 --- a/.gitignore +++ b/.gitignore @@ -26,3 +26,6 @@ dist/ /external /bazel-* htmlcov/ + +# AI +.claude/settings.local.json diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 4719d9b..35e398c 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -12,6 +12,15 @@ repos: - id: check-added-large-files args: ["--maxkb=1024"] - id: requirements-txt-fixer + - repo: https://github.com/astral-sh/ruff-pre-commit + # Ruff version. + rev: v0.15.8 + hooks: + # Run the linter. + - id: ruff-check + args: [--fix] + # Run the formatter. + - id: ruff-format - repo: https://github.com/codespell-project/codespell rev: v2.4.2 hooks: @@ -22,26 +31,3 @@ repos: rev: 1.4.3 hooks: - id: cmakelint - - repo: https://github.com/PyCQA/isort - rev: 8.0.1 - hooks: - - id: isort - args: - - '--profile=black' - - repo: https://github.com/psf/black - rev: 26.3.0 - hooks: - - id: black - args: - - '-vv' - - repo: https://github.com/PyCQA/isort - rev: 8.0.1 - hooks: - - id: isort - args: - - '--profile=black' - - repo: https://github.com/pycqa/flake8 - rev: 7.3.0 - hooks: - - id: flake8 - additional_dependencies: [Flake8-pyproject] diff --git a/MODULE.bazel b/MODULE.bazel index 163bb06..4c1f1cf 100644 --- a/MODULE.bazel +++ b/MODULE.bazel @@ -32,6 +32,14 @@ python = use_extension("@rules_python//python/extensions:python.bzl", "python") python.defaults(python_version = "3.12") python.toolchain(python_version = "3.12") +pip = use_extension("@rules_python//python/extensions:pip.bzl", "pip") +pip.parse( + hub_name = "vila_pip_deps", + python_version = "3.12", + requirements_lock = "//tests/vila:requirements_tvm_ffi_test.txt", +) +use_repo(pip, "vila_pip_deps") + wdk_config_ext = use_extension("@vila//vila/bazel/bzlmod:extensions.bzl", "wdk_configure_extension") use_repo(wdk_config_ext, "local_config_wdk") @@ -53,6 +61,10 @@ use_repo(hedron_ext, "hedron_compile_commands") sycl_config_ext = use_extension("@vila//vila/bazel/bzlmod:extensions.bzl", "sycl_configure_extension") use_repo(sycl_config_ext, "local_config_sycl") +# TVM FFI support (optional - requires: pip install apache-tvm-ffi) +tvm_ffi_ext = use_extension("@vila//vila/bazel/bzlmod:extensions.bzl", "tvm_ffi_extension") +use_repo(tvm_ffi_ext, "tvm_ffi") + register_toolchains( # to use sycl specify: --config=sycl "@local_config_sycl//:cc-toolchain-x64_sycl", diff --git a/WORKSPACE b/WORKSPACE index 9e6acd7..361fa02 100644 --- a/WORKSPACE +++ b/WORKSPACE @@ -26,7 +26,10 @@ vila_workspace1() load("@vila//vila:workspace2.bzl", vila_workspace2 = "workspace") -vila_workspace2(sycl = True) +vila_workspace2( + sycl = False, + tvm_ffi = True, +) # load("@bazel_features//:deps.bzl", "bazel_features_deps") # load("@bazel_skylib//:workspace.bzl", "bazel_skylib_workspace") @@ -47,3 +50,15 @@ python_register_toolchains( name = "local_config_python", python_version = "3.13", ) + +load("@rules_python//python:pip.bzl", "pip_parse") + +pip_parse( + name = "vila_pip_deps", + python_interpreter = "python", + requirements_lock = "//tests/vila:requirements_tvm_ffi_test.txt", +) + +load("@vila_pip_deps//:requirements.bzl", "install_deps") + +install_deps() diff --git a/pyproject.toml b/pyproject.toml index 6c9b77e..8d85ee1 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -47,3 +47,10 @@ disable = [ "R", "I", ] + +[tool.ruff] +line-length = 88 + +[tool.ruff.lint] +select = ["E", "F", "UP", "I"] +ignore = ["E203", "E231", "E241"] diff --git a/tests/vila/BUILD.bazel b/tests/vila/BUILD.bazel index 96ad83c..e25da92 100644 --- a/tests/vila/BUILD.bazel +++ b/tests/vila/BUILD.bazel @@ -14,8 +14,32 @@ See the License for the specific language governing permissions and limitations under the License. """ +load("@rules_cc//cc:defs.bzl", "cc_binary") +load("@rules_python//python:defs.bzl", "py_test") +load("@vila_pip_deps//:requirements.bzl", "requirement") load("//vila/bazel/vila:vila.bzl", "vila_cc_test") +cc_binary( + name = "tvm_ffi_call", + srcs = ["tvm_ffi_call.cpp"], + linkshared = True, + deps = [ + "@fmt", + "@tvm_ffi", + ], +) + +py_test( + name = "tvm_ffi_call_test", + srcs = ["tvm_ffi_call_test.py"], + data = [":tvm_ffi_call"], + deps = [ + requirement("apache-tvm-ffi"), + requirement("numpy"), + requirement("typing-extensions"), + ], +) + vila_cc_test( name = "config_test", srcs = ["config_test.cpp"], diff --git a/tests/vila/requirements_tvm_ffi_test.txt b/tests/vila/requirements_tvm_ffi_test.txt new file mode 100644 index 0000000..48c6717 --- /dev/null +++ b/tests/vila/requirements_tvm_ffi_test.txt @@ -0,0 +1,3 @@ +apache-tvm-ffi==0.1.9 +numpy==2.2.6 +typing-extensions==4.15.0 diff --git a/tests/vila/tvm_ffi_call.cpp b/tests/vila/tvm_ffi_call.cpp new file mode 100644 index 0000000..6745ee9 --- /dev/null +++ b/tests/vila/tvm_ffi_call.cpp @@ -0,0 +1,104 @@ +/** + * Copyright (C) 2026 The VILA Authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +// This file demonstrates how to use TVM FFI to export C++ functions +// that can be called from Python via tvm_ffi.load_module(). +// +// TVM FFI provides a simple mechanism to create shared libraries (.dll on Windows, +// .so on Linux) that can be loaded dynamically from Python. Functions exported +// using TVM_FFI_DLL_EXPORT_TYPED_FUNC are automatically discoverable. +// +// Key concepts: +// 1. TVM_FFI_DLL_EXPORT_TYPED_FUNC(ExportName, FunctionPtr) +// - ExportName: the name exposed to Python (e.g., "add_one_tensor") +// - FunctionPtr: a C++ function pointer with matching signature +// 2. The C++ function takes a tvm::ffi::Tensor and returns a tvm::ffi::Tensor +// 3. TVM FFI automatically converts numpy arrays to Tensor via DLPack +// +// Example usage in Python: +// import tvm_ffi +// import numpy as np +// mod = tvm_ffi.load_module("path/to/tvm_ffi_call.dll") +// func = mod.add_one_tensor +// x = np.array([1.0, 2.0, 3.0], dtype=np.float32) +// result = func(x) # Returns array([2.0, 3.0, 4.0], dtype=np.float32) + +#include + +#include +#include + +// CPU allocator for creating output tensors +struct CPUAlloc { + void AllocData(DLTensor* tensor) { + size_t size = tvm::ffi::GetDataSize(*tensor); + tensor->data = malloc(size); + } + + void FreeData(DLTensor* tensor) { + if (tensor->data != nullptr) { + free(tensor->data); + tensor->data = nullptr; + } + } +}; + +// AddOneTensor_ takes an input tensor and returns a new tensor with each +// element incremented by 1. This demonstrates: +// 1. Receiving a Tensor from Python (via DLPack conversion) +// 2. Creating a new output tensor using TVM FFI's tensor API +// 3. Performing element-wise operations on tensor data +// +// @param input The input tensor (float32) +// @return A new tensor with each element = input element + 1 +tvm::ffi::Tensor AddOneTensor_(const tvm::ffi::Tensor& input) { + // Get a view of the input tensor for reading + tvm::ffi::TensorView view(input); + + // Verify the input is float32 for this simple example + if (view.dtype().code != kDLFloat || view.dtype().bits != 32 || view.dtype().lanes != 1) { + TVM_FFI_THROW(ValueError) << "Expected float32 tensor, got dtype with code=" + << static_cast(view.dtype().code) + << ", bits=" << static_cast(view.dtype().bits); + } + + // Create output tensor with the same shape + tvm::ffi::Tensor output = tvm::ffi::Tensor::FromNDAlloc( + CPUAlloc(), // allocator + view.shape(), // shape + view.dtype(), // dtype (float32) + input.device() // device (cpu) + ); + + // Copy data and add 1 to each element + const float* in_data = static_cast(view.data_ptr()); + float* out_data = static_cast(output.data_ptr()); + + int64_t numel = view.numel(); + for (int64_t i = 0; i < numel; ++i) { + out_data[i] = in_data[i] + 1.0f; + } + + return output; +} + +// TVM_FFI_DLL_EXPORT_TYPED_FUNC registers the C++ function AddOneTensor_ as "add_one_tensor" +// in the shared library's exports. When Python loads this library via tvm_ffi, +// it can access the function as mod.add_one_tensor. +// +// The macro handles all the boilerplate for exporting a typed function to Python, +// including function signature registration and ABI compatibility. +TVM_FFI_DLL_EXPORT_TYPED_FUNC(add_one_tensor, AddOneTensor_) diff --git a/tests/vila/tvm_ffi_call_test.py b/tests/vila/tvm_ffi_call_test.py new file mode 100644 index 0000000..be920b0 --- /dev/null +++ b/tests/vila/tvm_ffi_call_test.py @@ -0,0 +1,84 @@ +# Copyright (C) 2026 The VILA Authors. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import os +from pathlib import Path + +import numpy as np +import tvm_ffi + + +def test_tvm_ffi_tensor(): + """Test that the TVM FFI C++ library handles tensor data correctly. + + This test loads the compiled TVM FFI shared library and verifies that: + 1. The library can be loaded successfully + 2. The add_one_tensor function processes 1D tensors correctly + 3. The add_one_tensor function processes 2D tensors correctly + """ + + # Determine the correct library extension for the platform + ext = ".dll" if os.name == "nt" else ".so" + prefix = "" if os.name == "nt" else "lib" + lib_name = f"{prefix}tvm_ffi_call{ext}" + + file_dir = Path(__file__).resolve().parent + workspace_root = file_dir.parents[1] + + # Prefer Bazel runfiles, then local workspace build outputs. + candidates = [ + Path(os.environ.get("TEST_SRCDIR", "")) / f"_main/tests/vila/{lib_name}", + Path(os.environ.get("RUNFILES_DIR", "")) / f"_main/tests/vila/{lib_name}", + file_dir / lib_name, + workspace_root / f"bazel-bin/tests/vila/{lib_name}", + Path.cwd() / lib_name, + ] + + lib_path = None + for candidate in candidates: + if candidate.exists(): + lib_path = str(candidate) + break + + if lib_path is None: + raise FileNotFoundError( + "Unable to locate TVM FFI shared library. Checked: " + + ", ".join(str(path) for path in candidates) + ) + + # Load the TVM FFI module + mod = tvm_ffi.load_module(lib_path) + func = mod.add_one_tensor + + # Test 1D tensor + x_1d = np.array([1.0, 2.0, 3.0], dtype=np.float32) + result_1d = np.from_dlpack(func(x_1d)) + expected_1d = np.array([2.0, 3.0, 4.0], dtype=np.float32) + assert np.allclose(result_1d, expected_1d), ( + f"1D: Expected {expected_1d}, got {result_1d}" + ) + + # Test 2D tensor + x_2d = np.array([[1.0, 2.0], [3.0, 4.0]], dtype=np.float32) + result_2d = np.from_dlpack(func(x_2d)) + expected_2d = np.array([[2.0, 3.0], [4.0, 5.0]], dtype=np.float32) + assert np.allclose(result_2d, expected_2d), ( + f"2D: Expected {expected_2d}, got {result_2d}" + ) + + print("TVM FFI tensor test passed!") + + +if __name__ == "__main__": + test_tvm_ffi_tensor() diff --git a/vila/bazel/bzlmod/extensions.bzl b/vila/bazel/bzlmod/extensions.bzl index a139249..e8a7a95 100644 --- a/vila/bazel/bzlmod/extensions.bzl +++ b/vila/bazel/bzlmod/extensions.bzl @@ -15,6 +15,7 @@ limitations under the License. """ load("@bazel_tools//tools/build_defs/repo:http.bzl", "http_archive") +load("@vila//vila/bazel:tvm_ffi_configure.bzl", "tvm_ffi_configure") load("@vila//vila/bazel/toolchains:bullseye_cc_configure.bzl", "bullseye_configure") load("@vila//vila/bazel/toolchains:sycl_cc_configure.bzl", "sycl_configure") load("@vila//vila/bazel/wdk:wdk_configure.bzl", "wdk_configure") @@ -90,3 +91,7 @@ bullseye_configure_extension = module_extension( sycl_configure_extension = module_extension( implementation = lambda ctx: sycl_configure(name = "local_config_sycl"), ) + +tvm_ffi_extension = module_extension( + implementation = lambda ctx: tvm_ffi_configure(name = "tvm_ffi"), +) diff --git a/vila/bazel/toolchains/sycl_cc_configure.bzl b/vila/bazel/toolchains/sycl_cc_configure.bzl index 12b9526..5cedce1 100644 --- a/vila/bazel/toolchains/sycl_cc_configure.bzl +++ b/vila/bazel/toolchains/sycl_cc_configure.bzl @@ -78,6 +78,22 @@ def get_llvm_version(repository_ctx): return v.basename return "llvm-unknown" +def _append_template_list_entries(existing_entries, new_entries): + r"""Append rendered list entries to a template field without leading commas. + + Args: + existing_entries: Existing rendered entries string (without surrounding []). + new_entries: List of quoted entries to append. + + Returns: + str: Rendered entries string safe for use inside [] in BUILD templates. + """ + existing = existing_entries.strip() + new_rendered = ",\n ".join(new_entries) + if not existing: + return " " + new_rendered + return existing + ",\n " + new_rendered + def _overwrite_sycl_msvc(repository_ctx, msvc_vars, target_arch): oneapi_path = find_oneapi_path(repository_ctx) llvm_version = get_llvm_version(repository_ctx) @@ -96,11 +112,14 @@ def _overwrite_sycl_msvc(repository_ctx, msvc_vars, target_arch): oneapi_path + "/lib/clang/%s/include" % llvm_version, oneapi_path + "/opt/compiler/include", ]) - msvc_vars["%{msvc_cxx_builtin_include_directories_" + target_arch + "}"] += ",\n " + ",\n ".join([ - "\"%s\"" % (oneapi_path + "/include"), - "\"%s\"" % (oneapi_path + "/lib/clang/%s/include" % llvm_version), - "\"%s\"" % (oneapi_path + "/opt/compiler/include"), - ]) + msvc_vars["%{msvc_cxx_builtin_include_directories_" + target_arch + "}"] = _append_template_list_entries( + msvc_vars.get("%{msvc_cxx_builtin_include_directories_" + target_arch + "}", ""), + [ + "\"%s\"" % (oneapi_path + "/include"), + "\"%s\"" % (oneapi_path + "/lib/clang/%s/include" % llvm_version), + "\"%s\"" % (oneapi_path + "/opt/compiler/include"), + ], + ) else: print("oneAPI DPC++ may not be installed, please check the environment variable ONEAPI_ROOT") diff --git a/vila/bazel/toolchains/sycl_cc_toolchain_config.bzl b/vila/bazel/toolchains/sycl_cc_toolchain_config.bzl index 8c56c7a..38f4e28 100644 --- a/vila/bazel/toolchains/sycl_cc_toolchain_config.bzl +++ b/vila/bazel/toolchains/sycl_cc_toolchain_config.bzl @@ -32,6 +32,7 @@ load( "variable_with_value", "with_feature_set", ) +load("@rules_cc//cc/common:cc_common.bzl", "cc_common") all_compile_actions = [ ACTION_NAMES.c_compile, @@ -84,6 +85,9 @@ all_link_actions = [ def _use_msvc_toolchain(ctx): return ctx.attr.cpu in ["x64_windows", "arm64_windows"] and (ctx.attr.compiler == "msvc-cl" or ctx.attr.compiler == "clang-cl") +def _filter_non_empty_paths(paths): + return [p for p in paths if p and p != "msvc_not_found"] + def _impl(ctx): if _use_msvc_toolchain(ctx): artifact_name_patterns = [ @@ -1430,7 +1434,7 @@ def _impl(ctx): features = features, action_configs = action_configs, artifact_name_patterns = artifact_name_patterns, - cxx_builtin_include_directories = ctx.attr.cxx_builtin_include_directories, + cxx_builtin_include_directories = _filter_non_empty_paths(ctx.attr.cxx_builtin_include_directories), toolchain_identifier = ctx.attr.toolchain_identifier, host_system_name = ctx.attr.host_system_name, target_system_name = ctx.attr.target_system_name, diff --git a/vila/bazel/tvm_ffi_configure.bzl b/vila/bazel/tvm_ffi_configure.bzl new file mode 100644 index 0000000..41f1d8d --- /dev/null +++ b/vila/bazel/tvm_ffi_configure.bzl @@ -0,0 +1,151 @@ +""" +Copyright (C) 2026 The VILA Authors. + +Locate TVM FFI from pip-installed apache-tvm-ffi package. + +TVM FFI (apache-tvm-ffi) is a pre-built binary package distributed via PyPI. +It provides C++ headers and a pre-built DLL/SO for calling TVM runtime from C++. + +This repository rule discovers the TVM FFI installation paths by running: + python -m tvm_ffi.config --includedir # Include directory + python -m tvm_ffi.config --dlpack-includedir # DLPack include + python -m tvm_ffi.config --libfiles # Library files (.lib on Windows) + +The rule then creates symlinks within the external repository to allow Bazel +to use these paths with relative includes (Bazel doesn't support absolute paths). + +Usage: + 1. Ensure apache-tvm-ffi is installed: pip install apache-tvm-ffi + 2. Add this repository rule to your WORKSPACE or MODULE.bazel + 3. Depend on @tvm_ffi//:tvm_ffi in your targets + +Example in MODULE.bazel: + tvm_ffi_ext = use_extension("@vila//vila/bazel/bzlmod:extensions.bzl", "tvm_ffi_extension") + use_repo(tvm_ffi_ext, "tvm_ffi") + +Example in WORKSPACE: + load("@vila//vila/bazel:tvm_ffi_configure.bzl", "tvm_ffi_configure") + tvm_ffi_configure(name = "tvm_ffi") +""" + +def _tvm_ffi_configure(repository_ctx): + """Repository rule implementation to locate TVM FFI from pip installation. + + This function is called by Bazel during the repository fetch phase. + It uses 'python -m tvm_ffi.config' to discover the TVM FFI installation + paths and creates a BUILD file that can be used by other targets. + + Args: + repository_ctx: The repository context provided by Bazel + + The function performs the following steps: + 1. Find Python executable (from attribute or Bazel toolchain) + 2. Query TVM FFI for include directories + 3. Query TVM FFI for library files + 4. Create symlinks in the external repository + 5. Generate a BUILD file with cc_library target + """ + + # Step 1: Find Python executable + # This is the Python that rules_python registers via python_register_toolchains + python_path = repository_ctx.which("python3") + if not python_path: + python_path = repository_ctx.which("python") + + # Verify Python was found + if not python_path or not python_path.exists: + fail( + "Python not found. Please either:\n" + + "1. Install apache-tvm-ffi: pip install apache-tvm-ffi\n" + + "2. Ensure python3 or python is available on PATH\n" + + "3. Use python_register_toolchains in your WORKSPACE/MODULE.bazel", + ) + + # Step 2: Get TVM FFI include directory + # The --includedir flag returns the path to tvm/ffi/include + result = repository_ctx.execute( + [str(python_path), "-m", "tvm_ffi.config", "--includedir"], + quiet = False, + ) + + # If TVM FFI is not installed, the command will fail + if result.return_code != 0: + fail("TVM FFI not found. Install with: pip install apache-tvm-ffi") + + # Store the include directory path (strip whitespace) + include_dir = result.stdout.strip() + + # Step 3: Get DLPack include directory + # DLPack is a standardization of in-memory tensor data structures + result = repository_ctx.execute( + [str(python_path), "-m", "tvm_ffi.config", "--dlpack-includedir"], + quiet = False, + ) + dlpack_dir = result.stdout.strip() + + # Step 4: Get library files for linking + # On Windows, this returns the .lib import library (not the .dll) + # On Linux, this would return .so files if available + result = repository_ctx.execute( + [str(python_path), "-m", "tvm_ffi.config", "--libfiles"], + quiet = False, + ) + + # Parse the comma-separated list of library files + lib_files = [f.strip() for f in result.stdout.strip().split(",") if f.strip()] + + # Step 5: Create symlinks in the external repository + # Bazel requires relative paths for includes, but TVM FFI uses absolute paths. + # We create symlinks from the external repo to the actual TVM FFI installation. + # This approach works across platforms and doesn't require copying files. + repository_ctx.symlink(include_dir, "include") + repository_ctx.symlink(dlpack_dir, "dlpack") + + # Step 6: Symlink the import library for linking + # On Windows, we need the .lib file to link against TVM FFI at compile time. + # The .dll is loaded at runtime via tvm_ffi.load_module() in Python. + lib_files_srcs = "" + if lib_files: + # Take the first library file (typically the import library) + for lib_file in lib_files: + # Create symlink to the .lib file + lib_name = repository_ctx.path(lib_file).basename + repository_ctx.symlink(lib_file, lib_name) + + # Format for BUILD file - use filegroup for proper handling + lib_files_srcs += '"%s",' % lib_name + + + # Step 7: Generate BUILD file + # The BUILD file defines a cc_library that propagates include paths. + # Using includes (not copts) ensures paths propagate to dependent targets. + # Note: TVM FFI is a pre-built binary, so we don't compile any source files. + # The library is loaded at runtime via dlopen/LoadLibrary from Python. + build_content = """ +package(default_visibility = ["//visibility:public"]) + +load("@rules_cc//cc:cc_library.bzl", "cc_library") + +cc_library( + name = "tvm_ffi", + # Relative paths via symlinked directories - Bazel requires this + includes = [ + "include", + "dlpack", + ], + # Link against TVM FFI import library (.lib on Windows) + # {lib_files_srcs} is intentionally left without leading comma if empty + srcs = [{lib_files_srcs}], + visibility = ["//visibility:public"], +) +""".format( + lib_files_srcs = lib_files_srcs, + ) + + # Write the generated BUILD file to the external repository root + repository_ctx.file("BUILD.bazel", build_content) + +# Define the repository rule for Bazel +# Repository rules are used to fetch and configure external dependencies. +# This rule is called once per Bazel invocation to set up the TVM FFI dependency. +tvm_ffi_configure = repository_rule(implementation = _tvm_ffi_configure) diff --git a/vila/workspace2.bzl b/vila/workspace2.bzl index 1110b26..e856e11 100644 --- a/vila/workspace2.bzl +++ b/vila/workspace2.bzl @@ -14,16 +14,18 @@ See the License for the specific language governing permissions and limitations under the License. """ +load("@vila//vila/bazel:tvm_ffi_configure.bzl", "tvm_ffi_configure") load("@vila//vila/bazel/toolchains:bullseye_cc_configure.bzl", "bullseye_configure") load("@vila//vila/bazel/toolchains:sycl_cc_configure.bzl", "sycl_configure") load("@vila//vila/bazel/wdk:wdk_configure.bzl", "wdk_configure") -def workspace(bullseye = False, sycl = False): +def workspace(bullseye = False, sycl = False, tvm_ffi = False): """Loads a set of vila dependencies. To be used in a WORKSPACE file. Args: bullseye: Whether to include Bullseye Coverage. sycl: Whether to include oneAPI DPC++ compiler. + tvm_ffi: Whether to include TVM FFI (requires: pip install apache-tvm-ffi). """ # Get Bullseye Coverage tool @@ -44,3 +46,7 @@ def workspace(bullseye = False, sycl = False): # Comment out the following line to use the bullseye. # "@local_config_bullseye//:cc-toolchain-x64_windows", ) + + # Get TVM FFI (optional, requires pip install apache-tvm-ffi) + if tvm_ffi: + tvm_ffi_configure(name = "tvm_ffi")