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
127 changes: 127 additions & 0 deletions test/distributed/test_transfer.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,127 @@
import gc
import unittest
import weakref
from unittest.mock import Mock

import torch
from torch.distributed._transfer import Endpoint
from torch.distributed._transfer._reference import MemoryBackend


class TestTransfer(unittest.TestCase):
def setUp(self):
self.backend = MemoryBackend()
self.a, self.b = Endpoint(self.backend), Endpoint(MemoryBackend())
self.src = torch.arange(8, dtype=torch.uint8)
self.dst = torch.zeros_like(self.src)
self.reg = self.a.register_tensor(self.src)
self.b.register_tensor(self.dst)
self.peer = self.a.connect(self.b.metadata())
self.local = self.a.prepare([(self.src.data_ptr(), 4, 0, 4, 2)], kind="host")
self.remote = self.a.prepare(
[(self.dst.data_ptr(), 4, 0, 4, 2)], kind="host", peer=self.peer
)

def tearDown(self):
self.a.close()
self.b.close()

def submit(self, operation="WRITE", **kwargs):
return self.a.submit(
operation, self.local, [1, 0], self.remote, [0, 1], **kwargs
)

def test_indexed_write_and_reuse(self):
for _ in range(2):
work = self.submit(notification=b"bucket-7")
work.wait()
self.assertEqual(self.dst.tolist(), [4, 5, 6, 7, 0, 1, 2, 3])
self.assertEqual(self.b.notifications(), {self.backend.name: [b"bucket-7"]})
self.assertEqual(self.b.notifications(), {})
work.close()

def test_read_and_standalone_notification(self):
self.dst.fill_(9)
work = self.submit("READ")
work.wait()
self.assertEqual(self.src.tolist(), [9] * 8)
work.close()
self.a.notify(self.peer, b"ack")
self.assertEqual(self.b.notifications(), {self.backend.name: [b"ack"]})

def test_close_in_dependency_order(self):
work = self.submit()
for close in (
self.reg.close,
self.local.close,
lambda: self.a.disconnect(self.peer),
):
with self.assertRaises(RuntimeError):
close()
work.close()
self.a.close()
self.a.close()
with self.assertRaises(RuntimeError):
self.a.register_tensor(self.src)

def test_timeout_and_release_failure_retain_resources(self):
self.backend.post = Mock(return_value="pending")
self.backend.poll = Mock(return_value="pending")
work = self.submit()
with self.assertRaises(TimeoutError):
work.wait(timeout=0)
with self.assertRaises(RuntimeError):
work.close()
self.backend.poll.return_value = "done"
self.backend.release_work = Mock(side_effect=RuntimeError("retry"))
with self.assertRaises(RuntimeError):
work.close()
with self.assertRaises(RuntimeError):
self.local.close()
self.backend.release_work.side_effect = None
work.close()

def test_post_exception_retains_handle_for_endpoint_cleanup(self):
self.backend.post = Mock(side_effect=RuntimeError("post uncertain"))
with self.assertRaises(RuntimeError):
self.submit()
with self.assertRaises(RuntimeError):
self.reg.close()
self.a.close() # reference provider reports terminal completion

def test_tensor_owner_is_retained(self):
owner = weakref.ref(self.src)
del self.src
gc.collect()
self.assertIsNotNone(owner())
self.a.close()
gc.collect()
self.assertIsNone(owner())

def test_host_wait_and_invalid_ordering(self):
event = Mock()
self.submit(ordering="host_wait", ready_event=event).close()
event.synchronize.assert_called_once_with()
for kwargs in (
{"ordering": "host_wait"},
{"ordering": "unknown"},
{"ready_event": event},
):
with self.assertRaises(ValueError):
self.submit(**kwargs)

def test_invalid_catalogs_indices_and_registration(self):
with self.assertRaises(ValueError):
self.a.prepare([(self.src.data_ptr() + 8, 4, 0)], kind="host")
with self.assertRaises(ValueError):
self.a.submit("WRITE", self.remote, [0], self.local, [0])
with self.assertRaises(ValueError):
self.a.submit("WRITE", self.local, [2], self.remote, [0])
with self.assertRaises(ValueError):
self.a.submit("WRITE", self.local, [], self.remote, [])
with self.assertRaises(ValueError):
self.a.register_tensor(torch.zeros(2, 2).t())


if __name__ == "__main__":
unittest.main()
54 changes: 54 additions & 0 deletions torch/distributed/_transfer/README.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,54 @@
# Transfer API — review MVP

Independent of ProcessGroup: move registered bytes between named peers using an
injected provider. No NIXL import/build dependency in Core. This private API is
experimental and deliberately not a production transport.

```python
import torch
from torch.distributed._transfer import Endpoint
from torch.distributed._transfer._reference import MemoryBackend

source, target = torch.arange(8, dtype=torch.uint8), torch.zeros(8, dtype=torch.uint8)
a, b = Endpoint(MemoryBackend()), Endpoint(MemoryBackend())
a.register_tensor(source)
b.register_tensor(target)
peer = a.connect(b.metadata()) # a real application exchanges this out of band
local = a.prepare([(source.data_ptr(), 4, 0, 4, 2)], kind="host")
remote = a.prepare([(target.data_ptr(), 4, 0, 4, 2)], kind="host", peer=peer)
work = a.submit("WRITE", local, [1, 0], remote, [0, 1])
work.wait()
work.close()
assert target.tolist() == [4, 5, 6, 7, 0, 1, 2, 3]
a.close()
b.close()
```

For NIXL, construct `Endpoint(nixl.torch_transfer.NixlBackend(native_agent))`.
The provider stays in NIXL. `MemoryBackend` is CPU-only, synchronous and in-process.

## Contract

- `register_tensor` retains a contiguous tensor. Raw `register` takes rows
`(address, bytes, device)` and a strong owner; the caller guarantees validity.
- `prepare` caches a descriptor catalog; optional `(stride, count)` fields keep
strided layouts compact. The provider validates indices/ranges/paired lengths.
- `submit` selects equally sized lists of indices. WRITE reads local memory;
READ writes local memory. Destination spans must not overlap.
- All methods serialize host access with one endpoint lock. GPU ordering is
separate: `caller_ready` trusts the caller; `host_wait` synchronizes the supplied
CUDA event before posting. No CUDA graph or stream-ordered completion support.
- `Work.poll()` returns `pending`, `done` or `error`. `wait()` timeout does not
cancel. Close succeeds only after `done`; pending/error work retains ownership.
- Close work, catalogs, peers, then registrations (or call endpoint.close).
Registration teardown conservatively waits for ALL endpoint catalogs to close.
- Stop remote access before teardown. Applications own rendezvous, authentication,
allocation lifetime, leases and acknowledgement protocols. Standalone `notify`
means local acceptance, not remote acknowledgement. Bytes remain opaque.

## Deliberately deferred

Active cancellation, failure recovery, revocation, metadata schemas, capability
negotiation, native dispatch helpers, slots, graph capture and performance parity.
No asynchronous Python interruption guarantee. Native GPU/end-to-end framework
validation is required before expanding this prototype or claiming speed parity.
9 changes: 9 additions & 0 deletions torch/distributed/_transfer/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,9 @@
"""Experimental, backend-independent point-to-point memory transfers.

This is a review prototype, not a stable or production-supported API.
"""

from ._api import Catalog, Endpoint, Registration, Work
from ._backend import Backend

__all__ = ["Backend", "Catalog", "Endpoint", "Registration", "Work"]
Loading