Skip to content
Closed
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
24 changes: 14 additions & 10 deletions cpp/include/tensorrt_llm/batch_manager/cacheTransceiver.h
Original file line number Diff line number Diff line change
Expand Up @@ -204,13 +204,15 @@ class BaseCacheTransceiver
{
public:
virtual ~BaseCacheTransceiver() = default;
virtual void respondAndSendAsync(LlmRequest* llmRequest) = 0;
// Transfers are asynchronous. Pass shared_ptr so the transceiver and its
// workers keep the request alive until the corresponding future resolves.
virtual void respondAndSendAsync(std::shared_ptr<LlmRequest> llmRequest) = 0;
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 +223,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 +254,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 +269,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 +278,10 @@ 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;
// Hold strong references while futures are outstanding so Python-side
// cleanup cannot leave C++ with dangling LlmRequest pointers.
std::vector<std::pair<std::shared_ptr<LlmRequest>, std::future<void>>> mSenderFutures;
std::vector<std::pair<std::shared_ptr<LlmRequest>, std::future<void>>> mRequesterFutures;
mpi::MpiComm const* mMpiWorldComm{nullptr};

std::shared_ptr<CacheTransceiverComm> mGroupComm;
Expand Down
10 changes: 10 additions & 0 deletions cpp/include/tensorrt_llm/executor/transferAgent.h
Original file line number Diff line number Diff line change
Expand Up @@ -288,6 +288,16 @@ class TransferStatus
virtual ~TransferStatus() = default;
[[nodiscard]] virtual bool isCompleted() const = 0;
virtual TransferState wait(int64_t timeout_ms = -1) const = 0;
/// Release the backend transfer request. If the request is still active,
/// backends may attempt to cancel it. A true return only means the backend
/// accepted release of the transfer handle. It is not proof that source or
/// destination memory is quiesced. Callers that release an in-progress
/// transfer must keep the affected memory ranges out of circulation unless
/// the backend provides a stronger completion guarantee.
[[nodiscard]] virtual bool release()
{
return false;
}
};

struct BaseAgentConfig
Expand Down
235 changes: 231 additions & 4 deletions cpp/tensorrt_llm/batch_manager/baseTransBuffer.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -21,11 +21,133 @@
#include "tensorrt_llm/common/logger.h"
#include "tensorrt_llm/common/opUtils.h"

#include <exception>
#include <mutex>

namespace tensorrt_llm::batch_manager
{

namespace
{

char const* bufferKindName(BufferKind kind)
{
switch (kind)
{
case BufferKind::kKV: return "kv";
case BufferKind::kKV_INDEXER: return "kv_indexer";
case BufferKind::kRNN: return "rnn";
}
return "unknown";
}

} // namespace

BufferIndexHolder::BufferIndexHolder(
BaseTransBufferManager* manager, Direction direction, std::optional<int> bufferId)
: mManager{manager}
, mDirection{direction}
, mBufferId{bufferId}
, mOwns{manager != nullptr}
{
}

BufferIndexHolder::~BufferIndexHolder()
{
reset();
}

BufferIndexHolder::BufferIndexHolder(BufferIndexHolder&& other) noexcept
: mManager{other.mManager}
, mDirection{other.mDirection}
, mBufferId{other.mBufferId}
, mOwns{other.mOwns}
{
other.mManager = nullptr;
other.mBufferId = std::nullopt;
other.mOwns = false;
}

BufferIndexHolder& BufferIndexHolder::operator=(BufferIndexHolder&& other) noexcept
{
if (this != &other)
{
reset();
mManager = other.mManager;
mDirection = other.mDirection;
mBufferId = other.mBufferId;
mOwns = other.mOwns;
other.mManager = nullptr;
other.mBufferId = std::nullopt;
other.mOwns = false;
}
return *this;
}

BufferIndexHolder BufferIndexHolder::acquireSend(BaseTransBufferManager& manager)
{
return BufferIndexHolder{&manager, Direction::kSend, manager.assignBufferIndexForSend()};
}

BufferIndexHolder BufferIndexHolder::acquireRecv(BaseTransBufferManager& manager)
{
return BufferIndexHolder{&manager, Direction::kRecv, manager.assignBufferIndexForRecv()};
}

void BufferIndexHolder::reset() noexcept
{
if (!mOwns || mManager == nullptr)
{
return;
}

try
{
if (mDirection == Direction::kSend)
{
mManager->freeBufferIndexForSend(mBufferId);
}
else
{
mManager->freeBufferIndexForRecv(mBufferId);
}
}
catch (std::exception const& e)
{
TLLM_LOG_ERROR(
"Exception while releasing cache transfer buffer index %d: %s", mBufferId.value_or(-1), e.what());
}
catch (...)
{
TLLM_LOG_ERROR("Unknown exception while releasing cache transfer buffer index %d", mBufferId.value_or(-1));
}

mManager = nullptr;
mBufferId = std::nullopt;
mOwns = false;
}

void BufferIndexHolder::poison() noexcept
{
if (!mOwns || mManager == nullptr)
{
return;
}

if (mDirection == Direction::kSend)
{
mManager->poisonBufferIndexForSend(mBufferId);
}
else
{
mManager->poisonBufferIndexForRecv(mBufferId);
}

mManager = nullptr;
mBufferId = std::nullopt;
mOwns = false;
}

BaseTransBufferManager::BaseTransBufferManager(
size_t transferBufferSize, nvinfer1::DataType dataType, std::optional<size_t> maxNumTokens)
: mDataType{dataType}
Expand Down Expand Up @@ -56,22 +178,58 @@ BaseTransBufferManager::BaseTransBufferManager(

std::optional<int> BaseTransBufferManager::assignBufferIndexForSend()
{
return assignBufferIndex(mConcurrenceSendResource, mSendBufferCount, mOnlyUseDynamicBuffer);
auto bufferId = assignBufferIndex(mConcurrenceSendResource, mSendBufferCount, mOnlyUseDynamicBuffer);
if (bufferId.has_value())
{
TLLM_LOG_DEBUG("Assigned send cache transfer buffer kind=%s index=%d outstanding=%d/%zu",
bufferKindName(getBufferKind()), bufferId.value(), mConcurrenceSendResource.mConcurrence.load(),
mSendBufferCount);
}
return bufferId;
}

void BaseTransBufferManager::freeBufferIndexForSend(std::optional<int> bufferId)
{
freeBufferIndex(mConcurrenceSendResource, bufferId, mSendBufferCount, mOnlyUseDynamicBuffer);
if (bufferId.has_value())
{
TLLM_LOG_DEBUG("Freed send cache transfer buffer kind=%s index=%d outstanding=%d/%zu",
bufferKindName(getBufferKind()), bufferId.value(), mConcurrenceSendResource.mConcurrence.load(),
mSendBufferCount);
}
}

void BaseTransBufferManager::poisonBufferIndexForSend(std::optional<int> bufferId) noexcept
{
poisonBufferIndex(mConcurrenceSendResource, bufferId, mSendBufferCount, mOnlyUseDynamicBuffer, "send");
}

std::optional<int> BaseTransBufferManager::assignBufferIndexForRecv()
{
return assignBufferIndex(mConcurrenceRecvResource, mRecvBufferCount, mOnlyUseDynamicBuffer);
auto bufferId = assignBufferIndex(mConcurrenceRecvResource, mRecvBufferCount, mOnlyUseDynamicBuffer);
if (bufferId.has_value())
{
TLLM_LOG_DEBUG("Assigned recv cache transfer buffer kind=%s index=%d outstanding=%d/%zu",
bufferKindName(getBufferKind()), bufferId.value(), mConcurrenceRecvResource.mConcurrence.load(),
mRecvBufferCount);
}
return bufferId;
}

void BaseTransBufferManager::freeBufferIndexForRecv(std::optional<int> bufferId)
{
freeBufferIndex(mConcurrenceRecvResource, bufferId, mRecvBufferCount, mOnlyUseDynamicBuffer);
if (bufferId.has_value())
{
TLLM_LOG_DEBUG("Freed recv cache transfer buffer kind=%s index=%d outstanding=%d/%zu",
bufferKindName(getBufferKind()), bufferId.value(), mConcurrenceRecvResource.mConcurrence.load(),
mRecvBufferCount);
}
}

void BaseTransBufferManager::poisonBufferIndexForRecv(std::optional<int> bufferId) noexcept
{
poisonBufferIndex(mConcurrenceRecvResource, bufferId, mRecvBufferCount, mOnlyUseDynamicBuffer, "recv");
}

std::tuple<std::vector<runtime::ITensor::SharedPtr>, size_t, bool> BaseTransBufferManager::getOrAllocateSendBuffers(
Expand Down Expand Up @@ -230,11 +388,23 @@ std::optional<int> BaseTransBufferManager::assignBufferIndex(
{
if (onlyUseDynamicBuffer)
{
TLLM_CHECK_WITH_INFO(!resource.mPoisoned.load(std::memory_order_relaxed),
"Cannot assign dynamic cache transfer buffer kind=%s because a previous transfer left dynamic transfer "
"memory poisoned. The process must restart before these memory ranges can be safely reused.",
bufferKindName(getBufferKind()));
return std::nullopt;
}
std::unique_lock lk(resource.mBuffersMutex);
resource.mBuffersCV.wait(
lk, [&resource, bufferCount]() { return static_cast<size_t>(resource.mConcurrence) < bufferCount; });
resource.mBuffersCV.wait(lk,
[&resource, bufferCount]()
{
return resource.mPoisoned.load(std::memory_order_relaxed)
|| static_cast<size_t>(resource.mConcurrence) < bufferCount;
});
TLLM_CHECK_WITH_INFO(!resource.mPoisoned.load(std::memory_order_relaxed),
"Cannot assign cache transfer buffer kind=%s because a previous transfer left the buffer pool poisoned. "
"The process must restart before these memory ranges can be safely reused.",
bufferKindName(getBufferKind()));
int bufferId = -1;
for (size_t i = 0; i < bufferCount; i++)
{
Expand Down Expand Up @@ -264,13 +434,70 @@ void BaseTransBufferManager::freeBufferIndex(
TLLM_CHECK(static_cast<size_t>(bufferId.value()) < bufferCount);
{
std::scoped_lock lk(resource.mBuffersMutex);
if (resource.mBufferIndexFlag[bufferId.value()] == 2)
{
TLLM_LOG_ERROR("Refusing to free poisoned cache transfer buffer kind=%s index=%d",
bufferKindName(getBufferKind()), bufferId.value());
return;
}
resource.mBufferIndexFlag[bufferId.value()] = 0;
}
resource.mConcurrence--;
resource.mBuffersCV.notify_one();
}
}

void BaseTransBufferManager::poisonBufferIndex(ConcurrenceResource& resource, std::optional<int> bufferId,
size_t bufferCount, bool onlyUseDynamicBuffer, char const* direction) noexcept
{
resource.mPoisoned.store(true, std::memory_order_relaxed);

if (onlyUseDynamicBuffer)
{
TLLM_LOG_ERROR(
"Poisoned dynamic %s cache transfer buffer kind=%s. Dynamic transfer memory cannot be safely reused; "
"the process must restart.",
direction, bufferKindName(getBufferKind()));
resource.mBuffersCV.notify_all();
return;
}

if (!bufferId.has_value())
{
TLLM_LOG_ERROR("Poisoned unknown %s cache transfer buffer kind=%s. The process must restart.", direction,
bufferKindName(getBufferKind()));
resource.mBuffersCV.notify_all();
return;
}

try
{
TLLM_CHECK(static_cast<size_t>(bufferId.value()) < bufferCount);
{
std::scoped_lock lk(resource.mBuffersMutex);
if (resource.mBufferIndexFlag[bufferId.value()] == 1)
{
resource.mBufferIndexFlag[bufferId.value()] = 2;
}
}
TLLM_LOG_ERROR(
"Poisoned %s cache transfer buffer kind=%s index=%d. The slot will not be returned to the pool because "
"transport quiescence is unknown; the process must restart before serving more transfers.",
direction, bufferKindName(getBufferKind()), bufferId.value());
}
catch (std::exception const& e)
{
TLLM_LOG_ERROR("Exception while poisoning %s cache transfer buffer kind=%s index=%d: %s", direction,
bufferKindName(getBufferKind()), bufferId.value_or(-1), e.what());
}
catch (...)
{
TLLM_LOG_ERROR("Unknown exception while poisoning %s cache transfer buffer kind=%s index=%d", direction,
bufferKindName(getBufferKind()), bufferId.value_or(-1));
}
resource.mBuffersCV.notify_all();
}

size_t BaseTransBufferManager::getRecvBufferCount()
{
return mRecvBufferCount;
Expand Down
Loading