Copyright (c) 2025-2025 Huawei Technologies Co., Ltd.
sysHAX-adapter 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.
Created: 2026-1-31
Desc: CPU inference cpu utils
*/
#include "cpu_utils.h"
#include <iostream>
#include <sstream>
#include <algorithm>
#include <cstdlib>
#include <cstdio>
#include <cstring>
#include <omp.h>
#include <sched.h>
#include <unistd.h>
void init_work_divider(WorkDivider *divider, int numas, int para, bool exp) {
if (numas <= 0){
std::cerr << "Error: numas must be greater than 0\n";
exit(1);
}
if (para <= 0){
std::cerr << "Error: para must be greater than 0\n";
exit(1);
}
divider->num_numas = exp ? 1 : numas;
divider->para = para;
divider->num_threads = omp_get_num_threads();
if (divider->num_threads % divider->para != 0){
std::cerr << "nthreads (" << divider->num_threads << ") %% para (" << divider->para << ") != 0\n";
exit(1);
}
divider->num_threads /= divider->para;
if (divider->num_threads % divider->num_numas != 0) {
std::cerr << "nthreads (" << divider->num_threads << ") %% numas (" << divider->num_numas << ") != 0\n";
exit(1);
}
divider->global_tid = omp_get_thread_num();
divider->threads_per_numa = divider->num_threads / divider->num_numas;
if (exp){
divider->tid = divider->global_tid % divider->threads_per_numa;
divider->my_numa = divider->global_tid / (divider->num_threads * para / numas);
divider->tid_in_numa = divider->tid;
}else{
divider->tid = divider->global_tid % divider->threads_per_numa + divider->global_tid / (divider->para * divider->threads_per_numa) * divider->threads_per_numa;
divider->my_numa = divider->tid / divider->threads_per_numa;
divider->tid_in_numa = divider->tid % divider->threads_per_numa;
}
}
void divide_work(const WorkDivider *divider, int work_per_thread, int work_remaining, SingleNumaWorkRange *pstSingleRange){
if (work_remaining == 0) {
pstSingleRange->begin_thread = divider->tid * work_per_thread;
pstSingleRange->end_thread = divider->tid * work_per_thread + work_per_thread;
pstSingleRange->work_per_thread = work_per_thread;
} else if (divider->tid < work_remaining) {
pstSingleRange->begin_thread = divider->tid * work_per_thread + divider->tid;
pstSingleRange->end_thread = (divider->tid + 1) * work_per_thread + (divider->tid + 1);
pstSingleRange->work_per_thread = work_per_thread + 1;
} else {
pstSingleRange->begin_thread = divider->tid * work_per_thread + work_remaining;
pstSingleRange->end_thread = (divider->tid + 1) * work_per_thread + work_remaining;
pstSingleRange->work_per_thread = work_per_thread;
}
}
void divide_all_work(const WorkDivider *divider, int total_workitems, SingleNumaWorkRange *pstSingleRange)
{
int work_per_thread = total_workitems / divider->num_threads;
int work_remaining = total_workitems % divider->num_threads;
divide_work(divider, work_per_thread, work_remaining, pstSingleRange);
}
void divide_work_first_numa(const WorkDivider *divider, int total_workitems, SingleNumaWorkRange *pstSingleRange)
{
if (divider->my_numa == 0) {
int work_per_thread = total_workitems / divider->threads_per_numa;
int work_remaining = total_workitems % divider->threads_per_numa;
divide_work(divider, work_per_thread, work_remaining, pstSingleRange);
return;
}
pstSingleRange->begin_thread = 0;
pstSingleRange->end_thread = 0;
pstSingleRange->work_per_thread = 0;
}
void divide_work_single_numa(const WorkDivider *divider, int total_workitems, SingleNumaWorkRange *pstSingleRange)
{
int workitem_per_thread = (total_workitems + divider->threads_per_numa - 1 ) / divider->threads_per_numa;
int begin_thread = workitem_per_thread * divider->tid_in_numa;
int end_thread = workitem_per_thread * (divider->tid_in_numa + 1);
if(end_thread > total_workitems){
end_thread = total_workitems;
}
pstSingleRange->begin_thread = begin_thread;
pstSingleRange->end_thread = end_thread;
pstSingleRange->work_per_thread = std::max(end_thread - begin_thread, 0);
}
void divide_work_all_numas(const WorkDivider *divider, int total_workitems, MultiNumaWorkRange *pstNulRange)
{
int max_workitems_per_numa = (total_workitems - 1) / divider->num_numas + 1;
int workitem_numa_begin = divider->num_numas == 1 ? 0 : divider->my_numa * max_workitems_per_numa;
int workitem_numa_end = workitem_numa_begin + max_workitems_per_numa;
if (workitem_numa_end > total_workitems) {
workitem_numa_end = total_workitems;
}
int workitems_my_numa = workitem_numa_end - workitem_numa_begin;
int max_workitems_per_thread = (workitems_my_numa - 1) / divider->threads_per_numa + 1;
int begin = divider->tid_in_numa * max_workitems_per_thread;
int end = begin + max_workitems_per_thread;
if (end > workitems_my_numa) {
end = workitems_my_numa;
}
pstNulRange->begin_numa = workitem_numa_begin;
pstNulRange->work_per_numa = max_workitems_per_numa;
pstNulRange->begin_thread = begin;
pstNulRange->end_thread = end;
pstNulRange->work_per_thread = end - begin;
}