/*
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);
}

// 分配单NUMA节点的工作
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;
}

// 在加载模型时,对于每个numa在其中分配工作
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);
}

// 分配所有NUMA节点的工作
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;
}