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");
}
}
};
#endif