Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 0 additions & 1 deletion build.sh
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
88 changes: 88 additions & 0 deletions build_custom_wheel.sh
Original file line number Diff line number Diff line change
@@ -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<timestamp>+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 "=============================================================================="

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

In https://github.com/google/tpu-sync/compare/main...wenjung2007:tpu-sync:wenjung/option-a-f0818d4?expand=1#diff-7b25c8e872c18ddf25a6241a6115017db172f1a2dc3536a97fee730612bf36abR125, can you change

echo "=== [7/7] 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_torch; print('TPU Sync torch module verified successfully!'); from tpu_sync.api.torch import kv_cache_manager; print('KV cache manager verified successfully!')"

to

echo "Installing generated wheel for verification..."
pip install --force-reinstall "${WHEEL_PATH}"

echo "Verifying Python import test from clean directory..."
(cd /tmp && python3 -c "import torch; print('PyTorch loaded:', torch.__version__); from tpu_sync.api.torch import torch_abi, kv_cache_manager; ext = torch_abi.load_extension('tpu_sync.frameworks.torch', '_tpu_raiden_torch'); print('TPU Sync torch module verified successfully:', ext); impl = kv_cache_manager._torch_impl(); print('KV cache manager verified successfully:', impl)")

36 changes: 15 additions & 21 deletions ci/build_wheel.sh
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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}"
Expand Down
26 changes: 25 additions & 1 deletion tpu_sync/api/torch/torch_abi.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 ``<package>.<stem>``, dispatching on torch ABI.

Expand All @@ -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.
Expand All @@ -111,8 +133,9 @@ def load_extension(package: str, stem: str):
# Version-suffixed variants take precedence over an unversioned <stem>.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_<stem>; the loader's module name
# (not the filename) determines the init symbol CPython looks up.
loader = importlib.machinery.ExtensionFileLoader(stem, str(path))
Expand All @@ -128,3 +151,4 @@ def load_extension(package: str, stem: str):
raise
setattr(pkg, stem, module)
return module

2 changes: 1 addition & 1 deletion tpu_sync/core/xla_compat.h
Original file line number Diff line number Diff line change
Expand Up @@ -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<RawBuffer>;

// Type-erased wrapper for CommonPjRtBuffer::ScopedHold.
Expand Down
8 changes: 4 additions & 4 deletions tpu_sync/frameworks/torch/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -25,23 +25,23 @@ 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",
],
)

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",
],
)

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",
],
)

Expand Down Expand Up @@ -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",
],
)

Expand Down
2 changes: 1 addition & 1 deletion tpu_sync/frameworks/torch/kv_cache_manager.h
Original file line number Diff line number Diff line change
Expand Up @@ -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"

Expand Down
2 changes: 1 addition & 1 deletion tpu_sync/frameworks/torch/torch_raw_transfer.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
2 changes: 1 addition & 1 deletion tpu_sync/frameworks/torch/torch_raw_transfer.h
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@
#include <vector>

#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"

Expand Down
24 changes: 11 additions & 13 deletions tpu_sync/frameworks/torch/torch_tpu_utils.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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()));
Expand Down Expand Up @@ -106,16 +112,8 @@ UnpackedTensor UnpackTorchTensor(const at::Tensor& tensor,
const size_t logical_slice_byte_size =
logical_physical_size / static_cast<size_t>(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(
Expand Down
2 changes: 1 addition & 1 deletion tpu_sync/frameworks/torch/torch_tpu_utils.h
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@
#include <vector>

#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"

Expand Down
2 changes: 1 addition & 1 deletion tpu_sync/frameworks/torch/torch_utils.h
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@
#include <vector>

#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"

Expand Down
2 changes: 1 addition & 1 deletion tpu_sync/frameworks/torch/weight_synchronizer.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
2 changes: 1 addition & 1 deletion tpu_sync/frameworks/torch/weight_synchronizer.h
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down