/*
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 model weight base class
*/

#ifndef MODEL_WEIGHT_BASE_H
#define MODEL_WEIGHT_BASE_H

#include <string>
#include <unordered_map>
#include <map>
#include <vector>
#include <torch/types.h>
#include "weight.h"
#include "tp_method.h"
#include "quantization/quantization_base.h"
#include "quantization/quantization_q4_0.h"
#include "quantization/quantization_q8_0.h"
#include "quantization/quantization_fp16.h"
#include "quantization/quantization_q8align.h"

namespace cpu_inference {

class ModelWeightBase
{
public:
    ModelWeightBase(const std::string& model_type, const std::string& quantization_type, 
        const std::map<int, std::vector<int>>& cpu_affinity);

    virtual bool add_weight(const std::string& weight_name, std::string dtype,
    std::vector<int> shape, void* data);
    virtual bool add_weight(const std::string& weight_name,
        const std::vector<int>& cur_numa_list,
        const std::vector<int>& cur_shape,
        std::string dtype, void* data);
    const Weight& get_weight(const std::string& weight_name) const;

    void set_model_type(std::string model_type);
    void set_quantization_type(std::string quantization_type);
    void set_cpu_affinity(std::map<int, std::vector<int>> cpu_affinity);

protected:
    std::string model_type;
    std::string quantization_type;
    std::unordered_map<std::string, StoreType> name_store_map;
    std::vector<int> numa_list;
    std::map<int, std::vector<int>> cpu_affinity;
    std::unordered_map<std::string, Weight> weight_map;

    void quantize_weight(const void* const src_data, void* dst_data, size_t element_cnt);

    std::vector<int> split_weight_row(const std::vector<int>& shape) const;
    template<typename BlockType>
    void quantize_weight_impl(const void* src_data, void* dst_data, size_t element_cnt, size_t QK);
};

template<typename BlockType>
void ModelWeightBase::quantize_weight_impl(const void* src_data, void* dst_data, size_t element_cnt, size_t QK)
{
    constexpr size_t quant_buffer_size = 256;
    assert(element_cnt % quant_buffer_size == 0);
    size_t block_num = element_cnt / quant_buffer_size;
    const size_t offset = quant_buffer_size / QK;
    for (size_t i = 0; i < block_num; ++i) {
        QuantizationBase<BlockType>::quantize(
            static_cast<const float16_t*>(src_data) + i * quant_buffer_size,
            static_cast<BlockType*>(dst_data) + i * offset,
            quant_buffer_size, "fp16");
    }
}

};  // namespace cpu_inference

#endif