Skip to content

cuda.core: experiment with torch 2.14+ stable PyObject->Tensor shim - #3033

Draft
leofang wants to merge 2 commits into
NVIDIA:mainfrom
leofang:experiment-stable-pyobj-shim
Draft

leofang wants to merge 2 commits into
NVIDIA:mainfrom
leofang:experiment-stable-pyobj-shim

Conversation

@leofang

@leofang leofang commented Oct 6, 2026

Copy link
Copy Markdown
Member

Experimental: adopt the stable PyObject -> Tensor C ABI shim that landed in PyTorch 2.14 (pytorch/pytorch#183323, merged as b69f7838a7) as the primary path in the tensor bridge, keeping the existing pyobj_to_aten_handle pointer-arithmetic trick as the fallback for torch 2.3-2.13.

Changes

  1. cuda_core/cuda/core/_memoryview.pyx: bump the outer version gate from (2, 12) to (2, 14). Resolves the BufferError: only CUDA Array Interface v3 or above is supported symptom I hit in Bump nightly PyTorch highest version to 2.14.1 #3031 against torch 2.14.1.
  2. cuda_core/cuda/core/_tensor_bridge.pyx: resolve torch_tensor_from_pyobject + aoti_torch_delete_tensor_object via dlsym / GetProcAddress on first call. Dispatch key is symbol presence, not version string:
    • resolved → stable shim path (owned handle, deleted after metadata reads via try/finally)
    • NULL → existing pointer-arithmetic fallback (correct for torch 2.3-2.13, cap-checked by _memoryview.pyx)

The outer version gate (2.3 <= torch <= 2.14) is unchanged. Once we drop torch < 2.14 support we can remove the upper bound entirely because the shim itself is a stable ABI contract.

Module-private diagnostics (used for the benchmark below):

  • _get_pyobj_path_counts() / _reset_pyobj_path_counts()
  • _get_shim_available()
  • _set_use_shim(bool) for A/B benchmarking against the same torch version

Verification (local)

Env: fresh conda, python 3.12, CUDA 13.4 toolkit, RTX 6000 Ada; test target cuda_core/tests/test_utils.py -k torch (63 tests: TestViewGPU[torch-*], TestViewCudaArrayInterfaceGPU[torch-*], test_torch_tensor_bridge_dtypes[*], test_ml_dtypes_bfloat16_torch_dlpack).

torch path result
2.14.1 shim 63 passed, 0 failed
2.13.0 fallback 63 passed, 0 failed

The ptr/itemsize assertions in test_torch_tensor_bridge_dtypes cover 13 dtypes end-to-end — if the shim's owned handle were disagreeing with torch's view of the storage, these would fail. Path counters confirm each call is routed as expected.

Benchmark

Same env, same tensor (32×32 float32 CUDA), 30 outer × 10 000 inner iterations, 2 000 warm-up, toggle via _set_use_shim:

      shim: best=1581.8ns median=1599.0ns stdev=8.1ns
  fallback: best=1393.4ns median=1400.9ns stdev=6.1ns

shim - fallback best-case: +188.4ns (+13.5%)

The ~190 ns overhead is the cost of the shim's at::Tensor copy-construction (bumps the storage refcount) plus the matching aoti_torch_delete_tensor_object. The pointer-arithmetic fallback does neither — the Python obj already keeps the storage alive via buf.exporting_obj = obj.

Open question — do we need the refcount bump?

The shim returns a new reference (// returns new reference in shim.h). For view_as_torch_tensor we only read metadata off the handle (get_data_ptr, get_sizes, get_strides, get_dtype, get_device_type, get_device_index) and the Python obj outlives the handle, so a borrowed variant that skipped the refcount bump would be functionally equivalent and would close the perf gap.

Worth asking upstream for a torch_tensor_borrow_from_pyobject (or equivalent flag on the existing shim)? Logging here for discussion.

-- Leo's bot

PyTorch >= 2.14 exposes torch_tensor_from_pyobject /
aoti_torch_delete_tensor_object as part of the stable C ABI (landed via
pytorch/pytorch#183323).  These let us obtain an AtenTensorHandle from a
Python tensor object without peering at THPVariable's internal layout.

Dispatch is driven by *symbol presence* via dlsym / GetProcAddress on
first call to view_as_torch_tensor:
  - resolved -> primary path (owned handle, deleted after metadata reads)
  - NULL     -> fall back to the existing pyobj_to_aten_handle
                pointer-arithmetic trick, which stays correct for
                torch 2.3-2.13 (the only versions where this matters).

The outer version gate in _memoryview.pyx (``2.3 <= torch <= 2.14``) is
unchanged -- it still bounds the fallback.  Once we drop support for
torch < 2.14 we can remove the upper bound entirely because the shim
itself is part of torch's stable ABI contract.

Diagnostics kept module-private:
  - _get_pyobj_path_counts() / _reset_pyobj_path_counts()
  - _get_shim_available()
  - _set_use_shim(bool) to force the fallback for A/B benchmarking.
@copy-pr-bot

copy-pr-bot Bot commented Oct 6, 2026

Copy link
Copy Markdown
Contributor

This pull request requires additional validation before any workflows can run on NVIDIA's runners.

Pull request vetters can view their responsibilities here.

Contributors can view more details about this message here.

@github-actions github-actions Bot added the cuda.core Everything related to the cuda.core module label Oct 6, 2026

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

cuda.core Everything related to the cuda.core module

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant