Skip to content

Commit 0631785

Browse files
arttianezhumeta-codesync[bot]
authored andcommitted
Add register_tensor/deregister_tensor to public TorchComm API (#2019)
Summary: Pull Request resolved: #2019 Add `register_address(void*, size_t)` and `deregister_address(void*)` as virtual methods on `TorchCommBackend` with no-op defaults. Add tensor-level `register_tensor`/`deregister_tensor` on `TorchComm` that delegates to the virtual methods. Bind as `register_tensor`/`deregister_tensor` in Python. ABI version bumped from 1.0 to 1.1 (new virtual methods on base class). Reviewed By: d4l3k Differential Revision: D100188720 fbshipit-source-id: 0c8ac5a2ae8b5137e8e36d86935a5083ddeead27
1 parent a4e7153 commit 0631785

8 files changed

Lines changed: 250 additions & 0 deletions

File tree

comms/torchcomms/TorchComm.cpp

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -512,6 +512,14 @@ int64_t TorchComm::get_device_transport() {
512512
return impl_->get_device_transport();
513513
}
514514

515+
void TorchComm::tensor_register(const at::Tensor& tensor) {
516+
impl_->tensor_register(tensor);
517+
}
518+
519+
void TorchComm::tensor_deregister(const at::Tensor& tensor) {
520+
impl_->tensor_deregister(tensor);
521+
}
522+
515523
// Communicator Management
516524
std::shared_ptr<TorchComm> TorchComm::split(
517525
const std::vector<int>& ranks,

comms/torchcomms/TorchComm.hpp

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -170,6 +170,10 @@ class TorchComm : public std::enable_shared_from_this<TorchComm> {
170170
// Throws if not supported by the backend.
171171
int64_t get_device_transport();
172172

173+
// Memory Registration API
174+
void tensor_register(const at::Tensor& tensor);
175+
void tensor_deregister(const at::Tensor& tensor);
176+
173177
std::shared_ptr<TorchCommBackend> getBackendImpl() const {
174178
return impl_;
175179
}

comms/torchcomms/TorchCommBackend.hpp

Lines changed: 32 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -361,6 +361,38 @@ class TorchCommBackend {
361361
std::string(getCommName()));
362362
}
363363

364+
/**
365+
* Register a tensor's memory with the backend for optimized data transfer.
366+
*
367+
* Pre-registers the memory region for zero-copy RDMA or similar transport
368+
* optimizations. Backends that support registration must override this.
369+
*
370+
* The caller is responsible for calling tensor_deregister() before the
371+
* tensor is freed. Failing to deregister leaks the backend registration
372+
* handle but does not cause a crash — the transport layer will clean up
373+
* on communicator finalization.
374+
*
375+
* @param tensor The tensor whose memory to register.
376+
*/
377+
virtual void tensor_register(const at::Tensor& /*tensor*/) {
378+
throw std::runtime_error(
379+
"[TorchCommBackend]: tensor_register not implemented for "
380+
"communicator:" +
381+
std::string(getCommName()));
382+
}
383+
384+
/**
385+
* Deregister a tensor's previously registered memory.
386+
*
387+
* @param tensor The tensor whose memory to deregister.
388+
*/
389+
virtual void tensor_deregister(const at::Tensor& /*tensor*/) {
390+
throw std::runtime_error(
391+
"[TorchCommBackend]: tensor_deregister not implemented for "
392+
"communicator:" +
393+
std::string(getCommName()));
394+
}
395+
364396
protected:
365397
void runAbortHooks() {
366398
for (const auto& [_, hook] : abortHooks_) {

comms/torchcomms/TorchCommPy.cpp

Lines changed: 31 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1446,6 +1446,37 @@ is valid until the communicator is destroyed.
14461446
RuntimeError: If the backend does not support device transport.
14471447
)",
14481448
py::call_guard<py::gil_scoped_release>())
1449+
.def(
1450+
"tensor_register",
1451+
[](TorchComm& self, const at::Tensor& tensor) {
1452+
self.tensor_register(tensor);
1453+
},
1454+
R"(
1455+
Register a tensor's memory with the communication backend.
1456+
1457+
Pre-registers the memory region for optimized data transfer (e.g.,
1458+
RDMA zero-copy). The caller must call tensor_deregister() before
1459+
freeing the tensor. Omitting deregistration leaks the backend handle
1460+
but does not crash; cleanup occurs on communicator finalization.
1461+
1462+
Args:
1463+
tensor: The tensor whose memory to register.
1464+
)",
1465+
py::arg("tensor"),
1466+
py::call_guard<py::gil_scoped_release>())
1467+
.def(
1468+
"tensor_deregister",
1469+
[](TorchComm& self, const at::Tensor& tensor) {
1470+
self.tensor_deregister(tensor);
1471+
},
1472+
R"(
1473+
Deregister a tensor's previously registered memory.
1474+
1475+
Args:
1476+
tensor: The tensor whose memory to deregister.
1477+
)",
1478+
py::arg("tensor"),
1479+
py::call_guard<py::gil_scoped_release>())
14491480

14501481
// Point-to-Point Operations
14511482
.def(

comms/torchcomms/_comms.pyi

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -707,6 +707,8 @@ class TorchComm:
707707
self, callback: Callable[[int, PostHookArgs], None]
708708
) -> RemovableHandle: ...
709709
def register_abort_hook(self, callback: Callable[[], None]) -> RemovableHandle: ...
710+
def tensor_register(self, tensor: Any) -> None: ...
711+
def tensor_deregister(self, tensor: Any) -> None: ...
710712

711713
def new_comm(
712714
backend: str,

comms/torchcomms/fake/TorchCommFake.hpp

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44

55
#include <comms/torchcomms/TorchCommBackend.hpp>
66
#include <comms/torchcomms/TorchWork.hpp>
7+
#include <unordered_set>
78
#include <vector>
89

910
namespace torch::comms {
@@ -193,6 +194,19 @@ class TorchCommFake : public TorchCommBackend {
193194
return abortEnabled_ && aborted_;
194195
}
195196

197+
// Memory registration — tracks registered addresses for test verification
198+
void tensor_register(const at::Tensor& tensor) override {
199+
registered_addrs_.insert(tensor.data_ptr());
200+
}
201+
202+
void tensor_deregister(const at::Tensor& tensor) override {
203+
registered_addrs_.erase(tensor.data_ptr());
204+
}
205+
206+
bool is_tensor_registered(const at::Tensor& tensor) const {
207+
return registered_addrs_.count(tensor.data_ptr()) > 0;
208+
}
209+
196210
private:
197211
bool initialized_;
198212
at::Device device_;
@@ -203,6 +217,7 @@ class TorchCommFake : public TorchCommBackend {
203217
bool abortEnabled_{false};
204218
bool aborted_{false};
205219
bool shouldFailReconfigure_{false};
220+
std::unordered_set<void*> registered_addrs_;
206221
};
207222

208223
} // namespace torch::comms
Lines changed: 79 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,79 @@
1+
#!/usr/bin/env python3
2+
# pyre-unsafe
3+
# Copyright (c) Meta Platforms, Inc. and affiliates.
4+
5+
import itertools
6+
import unittest
7+
8+
import torch
9+
from torchcomms.tests.integration.py.TorchCommTestHelpers import (
10+
skip_backend,
11+
TorchCommTestWrapper,
12+
)
13+
14+
15+
class RegisterTensorTest(unittest.TestCase):
16+
"""Test tensor_register/tensor_deregister via CPU broadcast.
17+
18+
Verifies that pre-registering a CPU tensor with the communicator
19+
enables zero-copy RDMA for CPU collectives. Backends that override
20+
tensor_register should inherit this test.
21+
"""
22+
23+
counts = [4, 1024, 1024 * 1024]
24+
dtypes = [torch.float, torch.int, torch.int8]
25+
26+
def setUp(self):
27+
self.wrapper = TorchCommTestWrapper()
28+
self.torchcomm = self.wrapper.get_torchcomm()
29+
self.rank = self.torchcomm.get_rank()
30+
self.num_ranks = self.torchcomm.get_size()
31+
32+
def tearDown(self):
33+
self.torchcomm = None
34+
self.wrapper = None
35+
36+
def _verify_broadcast_results(self, tensor, expected_value, count, dtype):
37+
expected = torch.ones(count, dtype=dtype, device="cpu") * expected_value
38+
if dtype == torch.float:
39+
self.assertTrue(
40+
torch.allclose(tensor, expected),
41+
f"CPU broadcast tensors not close enough for count={count}",
42+
)
43+
else:
44+
self.assertTrue(
45+
torch.equal(tensor, expected),
46+
f"CPU broadcast tensors not equal for count={count}",
47+
)
48+
49+
def _cpu_broadcast_with_tensor_register(self, count, dtype):
50+
"""Test CPU broadcast using TorchComm.tensor_register() public API."""
51+
root_rank = 0
52+
root_value = 99
53+
54+
if self.rank == root_rank:
55+
tensor = torch.ones(count, dtype=dtype, device="cpu") * root_value
56+
else:
57+
tensor = torch.zeros(count, dtype=dtype, device="cpu")
58+
59+
self.torchcomm.tensor_register(tensor)
60+
61+
try:
62+
work = self.torchcomm.broadcast(tensor, root_rank, False)
63+
self.assertTrue(work.is_completed())
64+
self._verify_broadcast_results(tensor, root_value, count, dtype)
65+
finally:
66+
self.torchcomm.tensor_deregister(tensor)
67+
68+
@skip_backend("nccl", msg="tensor_register not implemented for backend: ")
69+
@skip_backend("ncclx", msg="tensor_register not implemented for backend: ")
70+
@skip_backend("gloo", msg="tensor_register not implemented for backend: ")
71+
def test_cpu_broadcast_with_tensor_register(self):
72+
"""Test CPU broadcast with TorchComm.tensor_register() public API."""
73+
for count, dtype in itertools.product(self.counts, self.dtypes):
74+
with self.subTest(count=count, dtype=dtype):
75+
self._cpu_broadcast_with_tensor_register(count, dtype)
76+
77+
78+
if __name__ == "__main__":
79+
unittest.main()
Lines changed: 79 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,79 @@
1+
// Copyright (c) Meta Platforms, Inc. and affiliates.
2+
3+
#include <gtest/gtest.h>
4+
5+
#include <ATen/ATen.h>
6+
#include <comms/torchcomms/TorchComm.hpp>
7+
#include <comms/torchcomms/TorchCommFactory.hpp>
8+
#include <comms/torchcomms/fake/TorchCommFake.hpp>
9+
#include <cstdlib>
10+
11+
namespace torch::comms {
12+
13+
namespace {
14+
constexpr const char* kBackendName = "fake_test";
15+
constexpr const char* kBackendEnvKey = "TORCHCOMMS_BACKEND_LIB_PATH_FAKE_TEST";
16+
} // namespace
17+
18+
class TorchCommRegisterTensorTest : public ::testing::Test {
19+
protected:
20+
void SetUp() override {
21+
const char* lib_path = std::getenv("FAKE_TEST_BACKEND_LIB_PATH");
22+
ASSERT_NE(lib_path, nullptr) << "FAKE_TEST_BACKEND_LIB_PATH not set";
23+
setenv(kBackendEnvKey, lib_path, 1);
24+
25+
comm_ = new_comm(kBackendName, at::Device(at::kCPU), "register_test");
26+
ASSERT_NE(comm_, nullptr);
27+
}
28+
29+
void TearDown() override {
30+
comm_.reset();
31+
unsetenv(kBackendEnvKey);
32+
}
33+
34+
TorchCommFake* getFakeBackend() {
35+
return dynamic_cast<TorchCommFake*>(comm_->getBackendImpl().get());
36+
}
37+
38+
std::shared_ptr<TorchComm> comm_;
39+
};
40+
41+
TEST_F(TorchCommRegisterTensorTest, RegisterDelegatesToBackend) {
42+
auto tensor = at::ones({1024}, at::kFloat);
43+
auto* backend = getFakeBackend();
44+
ASSERT_NE(backend, nullptr);
45+
46+
EXPECT_FALSE(backend->is_tensor_registered(tensor));
47+
comm_->tensor_register(tensor);
48+
EXPECT_TRUE(backend->is_tensor_registered(tensor));
49+
}
50+
51+
TEST_F(TorchCommRegisterTensorTest, DeregisterRemovesRegistration) {
52+
auto tensor = at::ones({1024}, at::kFloat);
53+
auto* backend = getFakeBackend();
54+
ASSERT_NE(backend, nullptr);
55+
56+
comm_->tensor_register(tensor);
57+
EXPECT_TRUE(backend->is_tensor_registered(tensor));
58+
59+
comm_->tensor_deregister(tensor);
60+
EXPECT_FALSE(backend->is_tensor_registered(tensor));
61+
}
62+
63+
TEST_F(TorchCommRegisterTensorTest, RegisterMultipleTensors) {
64+
auto t1 = at::ones({1024}, at::kFloat);
65+
auto t2 = at::zeros({2048}, at::kByte);
66+
auto* backend = getFakeBackend();
67+
ASSERT_NE(backend, nullptr);
68+
69+
comm_->tensor_register(t1);
70+
comm_->tensor_register(t2);
71+
EXPECT_TRUE(backend->is_tensor_registered(t1));
72+
EXPECT_TRUE(backend->is_tensor_registered(t2));
73+
74+
comm_->tensor_deregister(t1);
75+
EXPECT_FALSE(backend->is_tensor_registered(t1));
76+
EXPECT_TRUE(backend->is_tensor_registered(t2));
77+
}
78+
79+
} // namespace torch::comms

0 commit comments

Comments
 (0)