Skip to content
Merged
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
29 changes: 19 additions & 10 deletions cpp/include/tensorrt_llm/batch_manager/cacheTransceiver.h
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,7 @@
#include <torch/custom_class.h>
#include <torch/python.h>
#include <type_traits>
#include <unordered_set>
#include <vector>

using SizeType32 = tensorrt_llm::runtime::SizeType32;
Expand Down Expand Up @@ -204,13 +205,15 @@ class BaseCacheTransceiver
{
public:
virtual ~BaseCacheTransceiver() = default;
virtual void respondAndSendAsync(LlmRequest* llmRequest) = 0;
// Async entry points take shared_ptr so the async worker holds a strong
// reference for the transfer lifetime; see CacheTransceiver::mSenderFutures.
virtual void respondAndSendAsync(std::shared_ptr<LlmRequest> llmRequest) = 0;
Comment thread
chienchunhung marked this conversation as resolved.
virtual void respondAndSendLayerWise(
RequestVector const& requests, std::shared_ptr<ContextProgress> const& progress)
= 0;

virtual void requestAndReceiveSync(LlmRequest* llmRequest) = 0;
virtual void requestAndReceiveAsync(LlmRequest* llmRequest) = 0;
virtual void requestAndReceiveSync(std::shared_ptr<LlmRequest> llmRequest) = 0;
virtual void requestAndReceiveAsync(std::shared_ptr<LlmRequest> llmRequest) = 0;

/// Check all requests transferring context, and return the requests that have completed or encountered an error.
virtual RequestStatuses checkContextTransferStatus(
Expand All @@ -221,7 +224,7 @@ class BaseCacheTransceiver

[[nodiscard]] virtual bool checkGenTransferComplete() const = 0;

virtual bool cancelRequest(LlmRequest* llmRequest) = 0;
virtual bool cancelRequest(std::shared_ptr<LlmRequest> llmRequest) = 0;
};

class CacheTransceiver : public BaseCacheTransceiver
Expand Down Expand Up @@ -252,13 +255,13 @@ class CacheTransceiver : public BaseCacheTransceiver

virtual ~CacheTransceiver();

void respondAndSendAsync(LlmRequest* llmRequest) override;
void respondAndSendAsync(std::shared_ptr<LlmRequest> llmRequest) override;

void respondAndSendLayerWise(
RequestVector const& requests, std::shared_ptr<ContextProgress> const& progress) override;

void requestAndReceiveSync(LlmRequest* llmRequest) override;
void requestAndReceiveAsync(LlmRequest* llmRequest) override;
void requestAndReceiveSync(std::shared_ptr<LlmRequest> llmRequest) override;
void requestAndReceiveAsync(std::shared_ptr<LlmRequest> llmRequest) override;

RequestStatuses checkContextTransferStatus(
std::optional<int> const& atLeastRequestNum = std::nullopt, bool markComplete = false) override;
Expand All @@ -267,7 +270,7 @@ class CacheTransceiver : public BaseCacheTransceiver

[[nodiscard]] bool checkGenTransferComplete() const override;

virtual bool cancelRequest(LlmRequest* llmRequest) override;
virtual bool cancelRequest(std::shared_ptr<LlmRequest> llmRequest) override;

private:
void initializeCommState();
Expand All @@ -276,8 +279,14 @@ class CacheTransceiver : public BaseCacheTransceiver

std::unique_ptr<CacheSender> mCacheSender;
std::unique_ptr<CacheReceiver> mCacheReceiver;
std::vector<std::pair<LlmRequest*, std::future<void>>> mSenderFutures;
std::vector<std::pair<LlmRequest*, std::future<void>>> mRequesterFutures;
// shared_ptr (not raw LlmRequest*) so the futures hold a strong reference for
// the transfer lifetime; otherwise Python's _terminate_request can drop the
// request while a C++ status check still dereferences it.
std::vector<std::pair<std::shared_ptr<LlmRequest>, std::future<void>>> mSenderFutures;
std::vector<std::pair<std::shared_ptr<LlmRequest>, std::future<void>>> mRequesterFutures;
// Dedup sets so observe-only timeout WARN logs fire at most once per stuck request.
std::unordered_set<LlmRequest::RequestIdType> mTimedOutSenderIds;
std::unordered_set<LlmRequest::RequestIdType> mTimedOutRequesterIds;
mpi::MpiComm const* mMpiWorldComm{nullptr};

std::shared_ptr<CacheTransceiverComm> mGroupComm;
Expand Down
24 changes: 24 additions & 0 deletions cpp/tensorrt_llm/batch_manager/baseTransBuffer.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,30 @@
namespace tensorrt_llm::batch_manager
{

void BufferIndexHolder::release() noexcept
{
if (!mHeld || mMgr == nullptr)
{
return;
}
try
{
if (mIsRecv)
{
mMgr->freeBufferIndexForRecv(mIndex);
}
else
{
mMgr->freeBufferIndexForSend(mIndex);
}
}
catch (...)
{
// noexcept: swallow so the destructor can never throw.
}
mHeld = false;
}

BaseTransBufferManager::BaseTransBufferManager(
size_t transferBufferSize, nvinfer1::DataType dataType, std::optional<size_t> maxNumTokens)
: mDataType{dataType}
Expand Down
81 changes: 81 additions & 0 deletions cpp/tensorrt_llm/batch_manager/baseTransBuffer.h
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,87 @@ enum class BufferKind : uint8_t
kRNN = 2
};

class BaseTransBufferManager;

/// @brief RAII holder for an index from BaseTransBufferManager::assignBufferIndexFor{Send,Recv}.
/// Releases on destruction (incl. exception unwind). Move-only; call release() on
/// the happy path or detach() when ownership is handed off downstream.
class BufferIndexHolder
Comment thread
chienchunhung marked this conversation as resolved.
{
public:
BufferIndexHolder() = default;

BufferIndexHolder(BaseTransBufferManager& mgr, std::optional<int> index, bool isRecv) noexcept
: mMgr(&mgr)
, mIndex(index)
, mHeld(index.has_value())
, mIsRecv(isRecv)
{
}

// Defined out-of-line in baseTransBuffer.cpp: release() calls
// BaseTransBufferManager methods, whose full definition appears later
// in this header.
~BufferIndexHolder()
{
release();
}

BufferIndexHolder(BufferIndexHolder const&) = delete;
BufferIndexHolder& operator=(BufferIndexHolder const&) = delete;

BufferIndexHolder(BufferIndexHolder&& other) noexcept
: mMgr(other.mMgr)
, mIndex(other.mIndex)
, mHeld(other.mHeld)
, mIsRecv(other.mIsRecv)
{
other.mHeld = false;
}

BufferIndexHolder& operator=(BufferIndexHolder&& other) noexcept
{
if (this != &other)
{
release();
mMgr = other.mMgr;
mIndex = other.mIndex;
mHeld = other.mHeld;
mIsRecv = other.mIsRecv;
other.mHeld = false;
}
return *this;
}

[[nodiscard]] std::optional<int> index() const noexcept
{
return mIndex;
}

[[nodiscard]] bool held() const noexcept
{
return mHeld;
}

/// @brief Relinquish ownership without releasing. Use when a downstream
/// owner (e.g. the formatter inside receiveSync) takes over the
/// release responsibility on the happy path.
std::optional<int> detach() noexcept
{
mHeld = false;
return mIndex;
}

/// @brief Release the slot now and disarm the destructor. Safe to call multiple times.
void release() noexcept;

private:
BaseTransBufferManager* mMgr{nullptr};
std::optional<int> mIndex{};
bool mHeld{false};
bool mIsRecv{true};
};

/// @brief Base class for cache transfer buffer management.
/// Handles buffer pool allocation, index assignment, and slicing.
/// Derived classes provide cache-specific size calculations.
Expand Down
10 changes: 5 additions & 5 deletions cpp/tensorrt_llm/batch_manager/cacheFormatter.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -523,6 +523,7 @@ void CacheFormatter::format(tensorrt_llm::batch_manager::TransferSession& sessio
// 5. send the buffer to the corresponding target. Ideally, we send only once (one buffer) for each target.

auto cacheBufferId = mCacheTransBufferManager->assignBufferIndexForSend();
BufferIndexHolder sendHolder(*mCacheTransBufferManager, cacheBufferId, /*isRecv=*/false);
int peerDuplicateHeadFactor = targetInfo.mPeerDupHeadFactor;
auto bufferTargetNum = targetNum / peerDuplicateHeadFactor;
auto ppRank = selfIdx
Expand Down Expand Up @@ -606,7 +607,7 @@ void CacheFormatter::format(tensorrt_llm::batch_manager::TransferSession& sessio

session.setTime(TransferSession::kTimeTransmissions);

mCacheTransBufferManager->freeBufferIndexForSend(cacheBufferId);
sendHolder.release();
session.setTime(TransferSession::kTimePostprocess);
}
TLLM_LOG_DEBUG(
Expand Down Expand Up @@ -849,6 +850,7 @@ void CacheFormatter::unformat(tensorrt_llm::batch_manager::TransferSession& sess
size_t remainNoCoverTargetNum = 0;
size_t bufferCoverTargetNum = 0;
std::optional<int> cacheBufferId = std::nullopt;
BufferIndexHolder recvHolder;
{
NVTX3_SCOPED_RANGE(formatInputAllocBuffer);

Expand All @@ -864,6 +866,7 @@ void CacheFormatter::unformat(tensorrt_llm::batch_manager::TransferSession& sess
{
cacheBufferId = mCacheTransBufferManager->assignBufferIndexForRecv();
}
recvHolder = BufferIndexHolder(*mCacheTransBufferManager, cacheBufferId, /*isRecv=*/true);
auto [recvSplitCachestmp, bufferCoverTargetNumtmp, onlyUseDynamicBuffer]
= mCacheTransBufferManager->getOrAllocateRecvBuffers(
cacheBufferId, static_cast<int>(targetNum), bufferEleSizes, bufferManager);
Expand Down Expand Up @@ -997,10 +1000,7 @@ void CacheFormatter::unformat(tensorrt_llm::batch_manager::TransferSession& sess
recvSplitCaches, outputBuffersPerWindow, destConfig, selfConfig, selfIdx, bufferManager);

bufferManager.getStream().synchronize();
if (cacheBufferId.has_value())
{
mCacheTransBufferManager->freeBufferIndexForRecv(cacheBufferId);
}
recvHolder.release();
}
session.setTime(TransferSession::kTimePostprocess);
}
Expand Down
Loading
Loading