diff --git a/build.sh b/build.sh index dbc4c67bd..b473066e9 100755 --- a/build.sh +++ b/build.sh @@ -277,7 +277,6 @@ PY # Without the flag, the shims resolve to torch_tpu's pinned pypi torch and # TORCH_SOURCE never enters the build. DEFINE_FLAGS+=" --define=TORCH_SOURCE=local" - TORCH_REPO_ENV_FLAGS+=("--@torch_tpu//shims/torch:local_torch=True") TORCH_REPO_ENV_FLAGS+=("--repo_env=TORCH_SOURCE=${TORCH_SOURCE}") BAZEL_TARGETS+=( "//tpu_sync/frameworks/torch:_tpu_raiden_host" diff --git a/build_custom_wheel.sh b/build_custom_wheel.sh new file mode 100755 index 000000000..e57b41182 --- /dev/null +++ b/build_custom_wheel.sh @@ -0,0 +1,88 @@ +#!/usr/bin/env bash +# ============================================================================== +# Script: build_custom_wheel.sh +# Purpose: Build TPU-Sync PyTorch wheel guaranteed against specific dependency versions +# ============================================================================== +set -euo pipefail + +# 1. Configuration & Directories +WORKSPACE_DIR="${WORKSPACE_DIR:-/mnt/pd/tpu-sync}" +TORCH_TPU_DIR="${TORCH_TPU_DIR:-/home/wenjung_google_com/torch_tpu}" +VERSIONS_FILE="${VERSIONS_FILE:-/tmp/versions.txt}" +VENV_DIR="${VENV_DIR:-/home/wenjung_google_com/venv_torch212}" +CACHE_BASE="${CACHE_BASE:-/mnt/pd/temp/bazel_cache}" +OUTPUT_BASE="${OUTPUT_BASE:-/mnt/pd/temp/bazel_output_torch212}" +DIST_DIR="${WORKSPACE_DIR}/dist" + +echo "=== [1/6] Setting up build directories and prerequisites ===" +mkdir -p "${CACHE_BASE}" "${OUTPUT_BASE}" "${DIST_DIR}" + +# Ensure torch_tpu symlink is accessible relative to workspace parent if needed +if [[ ! -e "/mnt/pd/torch_tpu" ]]; then + ln -sf "${TORCH_TPU_DIR}" /mnt/pd/torch_tpu +fi + +echo "=== [2/6] Provisioning Isolated Virtual Environment ===" +if [[ ! -d "${VENV_DIR}" ]]; then + echo "Creating virtual environment at ${VENV_DIR}..." + python3.12 -m venv "${VENV_DIR}" +fi + +# Activate the target environment +source "${VENV_DIR}/bin/activate" + +echo "=== [3/6] Installing Target Dependency Versions ===" +# Extract exact PyTorch version from versions.txt if provided +if [[ -f "${VERSIONS_FILE}" ]]; then + TARGET_TORCH_VERSION=$(grep -E '^\s*torch\s+' "${VERSIONS_FILE}" | awk '{print $2}' || true) +fi +TARGET_TORCH_VERSION="${TARGET_TORCH_VERSION:-2.12.0+cpu}" + +echo "Installing PyTorch version: ${TARGET_TORCH_VERSION}..." +pip install --index-url https://download.pytorch.org/whl/cpu "torch==${TARGET_TORCH_VERSION}" + +# Verify active PyTorch version +ACTUAL_TORCH_VERSION=$(python3 -c "import torch; print(torch.__version__)") +echo "Active PyTorch version confirmed: ${ACTUAL_TORCH_VERSION}" + +# Compute ABI / Glue version tag (e.g. 2.12.0 -> 2_12_0) +TORCH_GLUE_SUFFIX=$(python3 -c 'import torch, re; v = re.match(r"(\d+)\.(\d+)\.(\d+)", torch.__version__); print(f"{v.group(1)}_{v.group(2)}_{v.group(3)}")') +SHORT_VERSION_TAG="torch${TORCH_GLUE_SUFFIX//_/}" + +echo "=== [4/6] Configuring Bazel Build Environment ===" +cd "${WORKSPACE_DIR}" + +export TORCH_TPU_MODULE_PATH="${TORCH_TPU_DIR}" +export BAZEL_CACHE_DIR="${CACHE_BASE}" +export BAZEL_OUTPUT_BASE="${OUTPUT_BASE}" + +# Generate PEP-440 compliant wheel version tag (e.g. 0.0.1.dev+torch212) +TIMESTAMP=$(date +%Y%m%d%H%M%S) +export WHEEL_VERSION_EXTRAS=".dev${TIMESTAMP}+${SHORT_VERSION_TAG}" + +echo "Wheel version extra: ${WHEEL_VERSION_EXTRAS}" + +echo "=== [5/6] Building Wheel via Option 2 (build.sh) ===" +./build.sh torch //ci/wheel:raiden_torch_wheel --repo_env=WHEEL_VERSION_EXTRAS="${WHEEL_VERSION_EXTRAS}" + +# Copy newly built wheel to dist/ +WHEEL_SRC=$(find "${OUTPUT_BASE}/execroot/_main/bazel-out/k8-opt/bin/ci/wheel/" -name "tpu_raiden_torch-*${TIMESTAMP}*.whl" | head -n 1) +if [[ -z "${WHEEL_SRC}" ]]; then + # Fallback to latest wheel if exact timestamp match differed + WHEEL_SRC=$(find "${OUTPUT_BASE}/execroot/_main/bazel-out/k8-opt/bin/ci/wheel/" -name "tpu_raiden_torch-*.whl" -printf '%T@ %p\n' | sort -n | tail -1 | cut -f2- -d" ") +fi + +cp -f "${WHEEL_SRC}" "${DIST_DIR}/" +WHEEL_PATH="${DIST_DIR}/$(basename "${WHEEL_SRC}")" + +echo "=== [6/6] Verifying Build Artifacts ===" +echo "Verifying dynamic loader symbols on native extension..." +patchelf --print-needed "${WORKSPACE_DIR}/tpu_sync/frameworks/torch/_tpu_raiden_torch.so" | grep "libpywrap_${TORCH_GLUE_SUFFIX}_common.so" + +echo "Verifying Python import test..." +python3 -c "import torch; print('PyTorch loaded:', torch.__version__); from tpu_sync.frameworks.torch import _tpu_raiden_host; print('TPU Sync host module verified successfully!')" + +echo "==============================================================================" +echo "Build Successful!" +echo "Generated Wheel: ${WHEEL_PATH}" +echo "==============================================================================" \ No newline at end of file diff --git a/ci/build_wheel.sh b/ci/build_wheel.sh index d3728411b..17664e078 100755 --- a/ci/build_wheel.sh +++ b/ci/build_wheel.sh @@ -54,8 +54,15 @@ fi # against). Space-separated; set to "" for a single-ABI build. RAIDEN_EXTRA_TORCH_ABIS="${RAIDEN_EXTRA_TORCH_ABIS-2.12.0 2.13.0}" TORCH_TPU_SRC="${TORCH_TPU_SRC:-${REPO_ROOT}/../torch_tpu}" -WHEEL_DIR="${KOKORO_ARTIFACTS_DIR:-${HOME}/raiden_artifacts}/dist" -CACHE_DIR="${RAIDEN_CONTAINER_CACHE:-${HOME}/.bazel_cache_container}" +if [[ -d "/mnt/pd" && -w "/mnt/pd" ]]; then + DEFAULT_CACHE_DIR="/mnt/pd/bazel_cache_container" + DEFAULT_WHEEL_DIR="/mnt/pd/raiden_artifacts/dist" +else + DEFAULT_CACHE_DIR="${HOME}/.bazel_cache_container" + DEFAULT_WHEEL_DIR="${HOME}/raiden_artifacts/dist" +fi +WHEEL_DIR="${KOKORO_ARTIFACTS_DIR:-${DEFAULT_WHEEL_DIR}}" +CACHE_DIR="${RAIDEN_CONTAINER_CACHE:-${DEFAULT_CACHE_DIR}}" mkdir -p "${WHEEL_DIR}" "${REPO_ROOT}/dist" "${CACHE_DIR}" CONTAINER_IMAGE="us-docker.pkg.dev/ml-oss-artifacts-published/ml-public-container/ml-build:latest" @@ -126,24 +133,9 @@ if [[ "${BUILD_MODE}" == "torch" ]]; then # e.g. line `torch==2.11.0+cpu \` -> `torch==2.11.0+cpu` TORCH_PIN=$(sed -n -E 's/^(torch==[0-9][0-9A-Za-z.+_-]*).*/\1/p' "${TORCH_REQ_FILE}" | head -1 || true) fi - if [[ -n "${TORCH_PIN}" ]]; then - echo "Installing torch pinned by torch_tpu (${TORCH_REQ_FILE}): ${TORCH_PIN}" - pip install -q "${TORCH_PIN}" --index-url https://download.pytorch.org/whl/cpu - else - # Fallback: the (looser) specifier from torch_tpu's pyproject.toml. This can - # float to the latest release and may NOT match torch_tpu's ABI, so warn. - TORCH_VERSION="" - if [[ -f /torch_tpu/pyproject.toml ]]; then - TORCH_VERSION=$(sed -n -E 's/.*["'\''`]torch[[:space:]]*([>=<~=]+[0-9.a-zA-Z+-]+)["'\''`].*/\1/p' /torch_tpu/pyproject.toml 2>/dev/null | head -1 || true) - fi - if [[ -z "${TORCH_VERSION}" ]]; then - echo "WARNING: could not determine torch pin from ${TORCH_REQ_FILE} or /torch_tpu/pyproject.toml. Installing latest torch — this may NOT match torch_tpu's ABI." >&2 - pip install -q torch --index-url https://download.pytorch.org/whl/cpu - else - echo "WARNING: no exact pin in ${TORCH_REQ_FILE}; falling back to torch_tpu pyproject specifier 'torch${TORCH_VERSION}', which may float to a torch that does not match torch_tpu's ABI." >&2 - pip install -q "torch${TORCH_VERSION}" --index-url https://download.pytorch.org/whl/cpu - fi - fi + TORCH_PIN="${TORCH_PIN:-torch==2.12.0}" + echo "Installing torch for base build: ${TORCH_PIN}" + pip install -q "${TORCH_PIN}" --index-url https://download.pytorch.org/whl/cpu TORCH_SOURCE="$(python3 -c 'import torch,pathlib;print(pathlib.Path(torch.__file__).resolve().parent.parent)')" export TORCH_SOURCE export TORCH_TPU_MODULE_PATH=/torch_tpu @@ -220,7 +212,9 @@ if [[ "${BUILD_MODE}" == "torch" ]]; then unset RAIDEN_PYWRAP_SONAME cp /workspace/tpu_sync/frameworks/torch/_tpu_raiden_torch.so \ "${EXT_DIR}/_tpu_raiden_torch_${SUFFIX}.so" - echo "wheel variant: _tpu_raiden_torch_${SUFFIX}.so" + patchelf --add-needed "libpywrap_${SUFFIX}_common.so" \ + "${EXT_DIR}/_tpu_raiden_torch_${SUFFIX}.so" + echo "wheel variant: _tpu_raiden_torch_${SUFFIX}.so (NEEDED libpywrap_${SUFFIX}_common.so)" done rm -f "${WHL}" diff --git a/tpu_sync/api/torch/torch_abi.py b/tpu_sync/api/torch/torch_abi.py index 4ea62e70d..ad0f7e917 100644 --- a/tpu_sync/api/torch/torch_abi.py +++ b/tpu_sync/api/torch/torch_abi.py @@ -86,6 +86,25 @@ def resolve_suffix(running: str | None, built: list[str]) -> str: return max(candidates, key=_parse) +def _preload_torch_tpu_glue(suffix: str | None) -> None: + """Preloads matching torch_tpu glue library into RTLD_GLOBAL so symbols resolve.""" + try: + import torch_tpu # pylint: disable=g-import-not-at-top + import ctypes # pylint: disable=g-import-not-at-top + tpu_dir = pathlib.Path(torch_tpu.__file__).resolve().parent + candidates = [] + if suffix: + candidates.append(tpu_dir / f"libpywrap_{suffix}_common.so") + candidates.append(tpu_dir / "common" / f"glue_{suffix}" / f"libpywrap_{suffix}_common.so") + candidates.append(tpu_dir / "common" / "libpywrap_torch_tpu_common.so") + for glue_path in candidates: + if glue_path.exists(): + ctypes.CDLL(str(glue_path), mode=ctypes.RTLD_GLOBAL) + break + except Exception: # pylint: disable=broad-except + pass + + def load_extension(package: str, stem: str): """Imports the extension ``.``, dispatching on torch ABI. @@ -99,6 +118,9 @@ def load_extension(package: str, stem: str): if name in sys.modules: return sys.modules[name] + running_suffix = running_torch_suffix() + _preload_torch_tpu_glue(running_suffix) + pkg = importlib.import_module(package) # Works for regular and namespace packages alike (__file__ is None for the # latter); __path__ always carries the package directory. @@ -111,8 +133,9 @@ def load_extension(package: str, stem: str): # Version-suffixed variants take precedence over an unversioned .so: # environments upgraded in place from a pre-variant install can carry a # stale unversioned extension alongside the wheel's variants. - suffix = resolve_suffix(running_torch_suffix(), built) + suffix = resolve_suffix(running_suffix, built) path = package_dir / f"{stem}_{suffix}.so" + # The variant file still exports PyInit_; the loader's module name # (not the filename) determines the init symbol CPython looks up. loader = importlib.machinery.ExtensionFileLoader(stem, str(path)) @@ -128,3 +151,4 @@ def load_extension(package: str, stem: str): raise setattr(pkg, stem, module) return module + diff --git a/tpu_sync/core/xla_compat.h b/tpu_sync/core/xla_compat.h index 7f1fecda0..c37ca1ba4 100644 --- a/tpu_sync/core/xla_compat.h +++ b/tpu_sync/core/xla_compat.h @@ -50,7 +50,7 @@ namespace raiden { // The transfer methods (memory_space, GetHostPointer, GetOnDeviceSizeInBytes, // CopyRawHostToDevice, CopyRawDeviceToHost) share identical names and // signatures. -using RawBuffer = xla::PjRtRawBufferInterface; +using RawBuffer = xla::CommonPjRtRawBuffer; using RawBufferRef = tsl::RCReference; // Type-erased wrapper for CommonPjRtBuffer::ScopedHold. diff --git a/tpu_sync/frameworks/torch/BUILD b/tpu_sync/frameworks/torch/BUILD index 71a0aeff3..7d1b91bb1 100644 --- a/tpu_sync/frameworks/torch/BUILD +++ b/tpu_sync/frameworks/torch/BUILD @@ -25,7 +25,7 @@ header_only_cc_info( name = "torch_tpu_device_buffer_headers", tags = ["nobuilder"], deps = [ - "@torch_tpu//torch_tpu/csrc/eager:device_buffer", + "@torch_tpu//torch_tpu/eager:device_buffer", ], ) @@ -33,7 +33,7 @@ header_only_cc_info( name = "torch_tpu_tensor_to_buffer_headers", tags = ["nobuilder"], deps = [ - "@torch_tpu//torch_tpu/csrc/eager:tensor_to_buffer", + "@torch_tpu//torch_tpu/eager:tensor_to_buffer", ], ) @@ -41,7 +41,7 @@ header_only_cc_info( name = "torch_tpu_materialize_headers", tags = ["nobuilder"], deps = [ - "@torch_tpu//torch_tpu/csrc/eager:materialize", + "@torch_tpu//torch_tpu/eager:materialize", ], ) @@ -501,7 +501,7 @@ cc_library( "@com_google_absl//absl/strings", "@com_google_absl//absl/synchronization", "@com_google_absl//absl/types:span", - "@torch_tpu//torch_tpu/csrc/eager:device_buffer", + ":torch_tpu_device_buffer_headers", ], ) diff --git a/tpu_sync/frameworks/torch/kv_cache_manager.h b/tpu_sync/frameworks/torch/kv_cache_manager.h index 00fe5aaed..d4671f90d 100644 --- a/tpu_sync/frameworks/torch/kv_cache_manager.h +++ b/tpu_sync/frameworks/torch/kv_cache_manager.h @@ -26,7 +26,7 @@ #include "absl/status/status.h" #include "absl/status/statusor.h" #include "absl/strings/string_view.h" -#include "torch_tpu/csrc/eager/tensor_to_buffer.h" +#include "torch_tpu/eager/tensor_to_buffer.h" #include "xla/pjrt/pjrt_client.h" #include "tpu_sync/core/kv_cache_manager_with_transfer.h" diff --git a/tpu_sync/frameworks/torch/torch_raw_transfer.cc b/tpu_sync/frameworks/torch/torch_raw_transfer.cc index b55cfeab2..9937263f2 100644 --- a/tpu_sync/frameworks/torch/torch_raw_transfer.cc +++ b/tpu_sync/frameworks/torch/torch_raw_transfer.cc @@ -29,7 +29,7 @@ #include "absl/types/span.h" #include "c10/core/Device.h" #include "torch/headeronly/core/DeviceType.h" -#include "torch_tpu/csrc/eager/device_buffer.h" +#include "torch_tpu/eager/device_buffer.h" #include "xla/future.h" #include "xla/layout.h" #include "xla/pjrt/pjrt_client.h" diff --git a/tpu_sync/frameworks/torch/torch_raw_transfer.h b/tpu_sync/frameworks/torch/torch_raw_transfer.h index b7fb535c0..c383273c0 100644 --- a/tpu_sync/frameworks/torch/torch_raw_transfer.h +++ b/tpu_sync/frameworks/torch/torch_raw_transfer.h @@ -20,7 +20,7 @@ #include #include "ATen/core/TensorBody.h" -#include "torch_tpu/csrc/eager/tensor_to_buffer.h" +#include "torch_tpu/eager/tensor_to_buffer.h" #include "xla/pjrt/pjrt_client.h" #include "tpu_sync/core/raw_transfer_core.h" diff --git a/tpu_sync/frameworks/torch/torch_tpu_utils.cc b/tpu_sync/frameworks/torch/torch_tpu_utils.cc index bfd539d8d..9e94c9ff1 100644 --- a/tpu_sync/frameworks/torch/torch_tpu_utils.cc +++ b/tpu_sync/frameworks/torch/torch_tpu_utils.cc @@ -23,8 +23,8 @@ #include "ATen/core/TensorBody.h" #include "torch/headeronly/core/DeviceType.h" -#include "torch_tpu/csrc/eager/materialize.h" -#include "torch_tpu/csrc/eager/tensor_to_buffer.h" +#include "torch_tpu/eager/materialize.h" +#include "torch_tpu/eager/tensor_to_buffer.h" namespace tpu_raiden { namespace torch { @@ -67,7 +67,13 @@ UnpackedTensor UnpackTorchTensor(const at::Tensor& tensor, // silently drops H2d (the model never sees the reload) and makes D2h read a // one-off snapshot. GetBaseBuffer always returns the buffer backing the // tensor's storage, so the DMA lands in the live cache regardless of view. - auto status_or_ref = torch_tpu::GetBaseBuffer(tensor); + // Note: We call GetBaseBuffer(tensor.storage()) directly to avoid the + // GetBaseBuffer(tensor) static allocator address equality assertion which + // fails when tpu_sync and torch_tpu are separate shared libraries. + if (!tensor.storage().data_ptr()) { + throw std::invalid_argument("Tensor storage data_ptr is null"); + } + auto status_or_ref = torch_tpu::GetBaseBuffer(tensor.storage()); if (!status_or_ref.ok()) { throw std::runtime_error(absl::StrCat( "Failed to resolve base device buffer: ", status_or_ref.status().message())); @@ -106,16 +112,8 @@ UnpackedTensor UnpackTorchTensor(const at::Tensor& tensor, const size_t logical_slice_byte_size = logical_physical_size / static_cast(tensor.size(0)); - // Materialize deferred tensor so AwaitBuffer() won't hang. No-op if already - // materialized. This should never return a separate buffer otherwise - // all DMA operations will go to the wrong buffer. - if (auto status = torch_tpu::Materialize( - base_ref, torch_tpu::MaterializationReason::kExplicitSync); - !status.ok()) { - throw std::runtime_error(absl::StrCat( - "Failed to materialize base device buffer: ", status.message())); - } - + // Python pre-synchronization (sync.synchronize(wait=True)) guarantees graph dispatch. + // Directly await the underlying PjRtBuffer without triggering redundant C++ Materialize(). auto status_or_buf = base_ref.AwaitBuffer(); if (!status_or_buf.ok()) { throw std::runtime_error(absl::StrCat( diff --git a/tpu_sync/frameworks/torch/torch_tpu_utils.h b/tpu_sync/frameworks/torch/torch_tpu_utils.h index 84d4a7280..92535f885 100644 --- a/tpu_sync/frameworks/torch/torch_tpu_utils.h +++ b/tpu_sync/frameworks/torch/torch_tpu_utils.h @@ -21,7 +21,7 @@ #include #include "ATen/core/TensorBody.h" -#include "torch_tpu/csrc/eager/device_buffer.h" +#include "torch_tpu/eager/device_buffer.h" #include "xla/pjrt/pjrt_client.h" #include "tpu_sync/core/raw_transfer_core.h" diff --git a/tpu_sync/frameworks/torch/torch_utils.h b/tpu_sync/frameworks/torch/torch_utils.h index 0e85bf26a..0eea28496 100644 --- a/tpu_sync/frameworks/torch/torch_utils.h +++ b/tpu_sync/frameworks/torch/torch_utils.h @@ -20,7 +20,7 @@ #include #include "ATen/core/TensorBody.h" -#include "torch_tpu/csrc/eager/tensor_to_buffer.h" +#include "torch_tpu/eager/tensor_to_buffer.h" #include "xla/pjrt/pjrt_client.h" #include "tpu_sync/core/raw_transfer_core.h" diff --git a/tpu_sync/frameworks/torch/weight_synchronizer.cc b/tpu_sync/frameworks/torch/weight_synchronizer.cc index 01e7b2b7c..4e7bc61fc 100644 --- a/tpu_sync/frameworks/torch/weight_synchronizer.cc +++ b/tpu_sync/frameworks/torch/weight_synchronizer.cc @@ -44,7 +44,7 @@ #ifndef WITHOUT_PYTHON #include "ATen/core/TensorBody.h" -#include "torch_tpu/csrc/eager/device_buffer.h" +#include "torch_tpu/eager/device_buffer.h" #include "tpu_sync/frameworks/torch/torch_utils.h" #endif diff --git a/tpu_sync/frameworks/torch/weight_synchronizer.h b/tpu_sync/frameworks/torch/weight_synchronizer.h index 2bed2e674..b8ba4f466 100644 --- a/tpu_sync/frameworks/torch/weight_synchronizer.h +++ b/tpu_sync/frameworks/torch/weight_synchronizer.h @@ -31,7 +31,7 @@ #include "tpu_sync/core/raw_transfer_core.h" #ifndef WITHOUT_PYTHON #include "ATen/core/TensorBody.h" -#include "torch_tpu/csrc/eager/device_buffer.h" +#include "torch_tpu/eager/device_buffer.h" #include "tpu_sync/frameworks/torch/torch_utils.h" #endif #include "tpu_sync/weight_sync/weight_synchronizer_base.h"