* Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved.
* MindIE is licensed under Mulan PSL v2.
* You can use this software according to the terms and conditions of the Mulan PSL v2.
* You may obtain a copy of Mulan PSL v2 at:
* http://license.coscl.org.cn/MulanPSL2
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND,
* EITHER EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT,
* MERCHANTABILITY OR FIT FOR A PARTICULAR PURPOSE.
* See the Mulan PSL v2 for more details.
*/
#include "load_balancer.h"
#include "log.h"
namespace mindie_llm {
LoadBalancer::LoadBalancer(const std::vector<EnginePerDPSPtr> &enginePerDPs, size_t waveNumPerDP, size_t thresholdPerDP,
size_t intervalMs)
: enginePerDPs_(enginePerDPs), dpSize_(enginePerDPs.size()), waveNumPerDP_(waveNumPerDP),
thresholdPerDP_(thresholdPerDP), waveId_(0), intervalMs_(intervalMs), timer_(io_)
{
balancePolicyPtr_ = MakeBalancePolicy(BalancePolicyType::ROUND_ROBIN);
if (balancePolicyPtr_ == nullptr) {
throw std::runtime_error("LoadBalancer::LoadBalancer: Failed to create BalancePolicy");
}
lastIntervalTime_ = intervalMs_;
timer_.expires_after(std::chrono::milliseconds(intervalMs_));
timer_.async_wait(boost::bind(&LoadBalancer::ExecutePeriodicTask, this, boost::asio::placeholders::error));
worker_ = std::thread([this]() { io_.run(); });
pthread_setname_np(worker_.native_handle(), "LoadBalancer");
}
LoadBalancer::~LoadBalancer()
{
Stop();
if (worker_.joinable()) {
worker_.join();
}
}
void LoadBalancer::AddSeqGroup(SequenceGroupSPtr &seqGroup)
{
if (seqGroup == nullptr) {
MINDIE_LLM_LOG_ERROR("Attempt to add a null SequenceGroup.");
return;
}
if (stop_) {
MINDIE_LLM_LOG_INFO("LoadBalancer stopped.");
return;
}
seqGroups_.PushBack(seqGroup);
if (seqGroups_.Size() >= waveNumPerDP_) {
boost::asio::post(io_, boost::bind(&LoadBalancer::TriggerImmediatelyTask, this));
}
}
void LoadBalancer::TriggerImmediatelyTask()
{
timer_.cancel();
MINDIE_LLM_LOG_DEBUG("TriggerImmediatelyTask.");
Distribute();
timer_.expires_after(std::chrono::milliseconds(intervalMs_));
timer_.async_wait(boost::bind(&LoadBalancer::ExecutePeriodicTask, this, boost::asio::placeholders::error));
}
void LoadBalancer::ExecutePeriodicTask(const boost::system::error_code &ec)
{
if (ec == boost::asio::error::operation_aborted || (stop_ && seqGroups_.Empty())) {
return;
}
static int times = 0;
if ((times++ % 3000) == 0) {
MINDIE_LLM_LOG_DEBUG("ExecutePeriodicTask.");
}
Distribute();
timer_.expires_after(std::chrono::milliseconds(intervalMs_));
timer_.async_wait(boost::bind(&LoadBalancer::ExecutePeriodicTask, this, boost::asio::placeholders::error));
}
void LoadBalancer::Stop()
{
stop_ = true;
io_.stop();
MINDIE_LLM_LOG_INFO("LoadBalancer stopped.");
}
void LoadBalancer::Distribute()
{
if (seqGroups_.Empty() && candidates_.empty()) {
return;
}
std::vector<size_t> dpRankIds = GetDistributedDpRankIds();
if (dpRankIds.empty()) {
MINDIE_LLM_LOG_WARN("There is not scheduler need to distribute sequence group.");
return;
}
PrepCandidates(dpRankIds.size(), waveNumPerDP_);
DistributeSeqGroups(dpRankIds);
}
std::vector<size_t> LoadBalancer::GetDistributedDpRankIds()
{
std::vector<size_t> dpRankIds;
uint64_t minFreeBlockNum = std::numeric_limits<uint64_t>::max();
for (size_t i = 0; i < enginePerDPs_.size(); ++i) {
SchedulerMetric metric = enginePerDPs_[i]->scheduler->CollectSchedulerMetric();
dpRankIds.push_back(i);
uint64_t freeBlockNum = metric.blockInfo.freeNpuBlockNum_;
minFreeBlockNum = std::min(minFreeBlockNum, freeBlockNum);
}
BalancerConstraintParam param(waveNumPerDP_, dpRankIds.size(), minFreeBlockNum);
balancePolicyPtr_->SetConstraint(param);
return dpRankIds;
}
void LoadBalancer::PrepCandidates(size_t numDp, size_t maxNumPerDp)
{
SequenceGroupSPtr seqGroup;
while (seqGroups_.PopFront(seqGroup)) {
candidates_.push_back(seqGroup);
if (candidates_.size() >= maxNumPerDp * numDp) {
break;
}
}
}
void LoadBalancer::DistributeSeqGroups(const std::vector<size_t> &dpRankIds)
{
if (this->candidates_.empty()) {
return;
}
waveId_++;
std::vector<std::vector<SequenceGroupSPtr>> candidatesForScheds = balancePolicyPtr_->Apply(this->candidates_);
for (size_t i = 0; i < candidatesForScheds.size(); i++) {
auto candidatesForSched = candidatesForScheds[i];
auto &enginePerDP = enginePerDPs_[dpRankIds[i]];
for (auto seqGroupSPtr : candidatesForSched) {
if (seqGroupSPtr == nullptr) {
MINDIE_LLM_LOG_ERROR("null seqGroup in candidates");
continue;
}
seqGroupSPtr->waveId_ = waveId_;
enginePerDP->scheduler->AddSeqGroup(seqGroupSPtr);
MINDIE_LLM_LOG_INFO_REQUEST("[LlmEngine|Request-Enter Waiting] Load Balance dispatch request to DP RankId: "
<< dpRankIds[i] << ", requestId: " << seqGroupSPtr->metrics_.inferReqId_
<< "seqId: " << seqGroupSPtr->firstSeq->seqId_ << ").");
enginePerDP->addedRequestNum++;
}
}
}
LoadBalancerPtr MakeLoadBalancer(const std::vector<std::shared_ptr<EnginePerDP>> &enginePerDPs, size_t waveNumPerDP,
size_t thresholdPerDP, size_t intervalMs)
{
return std::make_unique<LoadBalancer>(enginePerDPs, waveNumPerDP, thresholdPerDP, intervalMs);
}
}