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

如何创建取值仅为1或-1的可学习参数/权重向量?

实现仅取1或-1的可学习乘法权重

完全可行,但不能直接将权重约束为离散的1/-1——毕竟梯度下降依赖连续参数来计算梯度。我们可以通过连续参数加近似离散的激活映射来实现:训练时用接近±1的连续值保证梯度可导,推理时切换为严格的1/-1。

修改后的代码实现

from tensorflow.keras.layers import Layer
from tensorflow.keras.layers import Input
from tensorflow.keras.models import Model
import tensorflow as tf

class BinaryLearnableMultiplier(Layer):
    def __init__(self, **kwargs):
        super(BinaryLearnableMultiplier, self).__init__(**kwargs)

    def build(self, input_shape):
        # 初始化权重为[-1,1]范围内的连续值
        self.kernel = self.add_weight(name='kernel',
                                      shape=(input_shape[-1],),
                                      initializer='uniform',  # 均匀初始化在[-1,1]区间
                                      trainable=True)
        super(BinaryLearnableMultiplier, self).build(input_shape)

    def call(self, inputs, training=None):
        if training:
            # 训练阶段:用tanh放大权重,让输出接近±1但保持连续(保证梯度可计算)
            scaled_kernel = tf.tanh(self.kernel * 10)  # 系数10越大,输出越接近±1,但梯度更陡峭
        else:
            # 推理阶段:转成严格的1或-1
            scaled_kernel = tf.sign(self.kernel)
            # 处理权重为0的极端情况(比如初始化时可能出现),默认转为1
            scaled_kernel = tf.where(scaled_kernel == 0, tf.ones_like(scaled_kernel), scaled_kernel)
        return inputs * scaled_kernel

# 构建测试模型
inputs = Input(shape=(64,))
multiplier = BinaryLearnableMultiplier()(inputs)
model = Model(inputs=inputs, outputs=multiplier)

关键逻辑说明

  • 训练阶段:tf.tanh(self.kernel * 10)把连续权重映射到接近±1的区间,既保留了梯度可导性,又让权重的实际作用接近目标的1/-1。系数10可灵活调整:数值越大,输出越接近离散值,但梯度会更陡峭;数值越小,梯度更稳定,但离散性稍弱。
  • 推理阶段:用tf.sign直接将权重转为严格的1或-1,同时处理可能出现的0值(避免乘0导致信息丢失)。
  • 初始化:改用uniform初始化在[-1,1]区间,让初始值更接近目标的±1,加快模型收敛。

另一种可选实现(用sigmoid映射):

# 训练阶段替代方案:将权重转成0/1后再映射为1/-1
scaled_kernel = 2 * tf.sigmoid(self.kernel * 10) - 1

这个逻辑和tanh方案效果类似,同样能实现接近±1的连续输出。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.30 03:45:08