You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何实现支持中间类权重分配的one-hot编码方法?

相邻类别插值软编码实现

你需要的这种编码不属于常规独热编码范畴,是针对有序类别设计的、支持中间值线性分配权重的软编码,没有开箱即用的官方编码器,手写实现逻辑非常简单。

需求规则梳理

现有3个有序类别0, 1, 2,编码规则为:

  • 若输入值为整数类标(如0/1/2),输出和普通独热编码一致,对应位置为1,其余为0
  • 若输入值落在两个相邻类别之间,权重按距离线性分配给左右两个类别,其余位置为0:
    • 输入0.4,距离0类0.4、距离1类0.6,输出[0.4, 0.6, 0]
    • 输入1.8,距离1类0.2、距离2类0.8,输出[0, 0.2, 0.8]

常规pandas独热编码实现效果参考:
pandas普通独热编码示例

代码实现

以下是基于numpy的批量编码实现,支持单值、列表、numpy数组输入:

import numpy as np

def interp_soft_encode(values, n_classes):
    """
    有序类别中间值线性插值软编码
    :param values: 待编码的一维数值序列,取值范围需落在[0, n_classes-1]区间内
    :param n_classes: 总类别数量
    :return: shape为(样本数, n_classes)的编码结果数组
    """
    values = np.asarray(values, dtype=np.float64).flatten()
    encode_res = np.zeros((values.shape[0], n_classes), dtype=np.float64)

    floor_idx = np.floor(values).astype(int)
    frac_weight = values - floor_idx

    # 处理刚好落在整数类上的样本
    exact_mask = frac_weight == 0
    encode_res[exact_mask, floor_idx[exact_mask]] = 1

    # 处理落在两个类别之间的样本
    interp_mask = ~exact_mask
    valid_floor = floor_idx[interp_mask]
    valid_frac = frac_weight[interp_mask]
    encode_res[interp_mask, valid_floor] = 1 - valid_frac
    encode_res[interp_mask, valid_floor + 1] = valid_frac

    return encode_res

效果测试

test_input = [0.4, 1.8, 0, 2, 1.2]
print(interp_soft_encode(test_input, n_classes=3))

运行输出:

[[0.4 0.6 0. ]
 [0.  0.2 0.8]
 [1.  0.  0. ]
 [0.  0.  1. ]
 [0.  0.8 0.2]]

和需求给出的预期结果完全一致。

适配sklearn工作流

如果需要接入sklearn的Pipeline、和其他预处理组件串联使用,可以简单封装为标准转换器:

from sklearn.base import BaseEstimator, TransformerMixin

class InterpSoftEncoder(BaseEstimator, TransformerMixin):
    def __init__(self, n_classes):
        self.n_classes = n_classes

    def fit(self, X, y=None):
        return self

    def transform(self, X):
        return interp_soft_encode(X, self.n_classes)

内容的提问来源于stack exchange,提问作者debrises

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.01 18:42:33