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

如何在tf.keras中创建带可训练稀疏权重矩阵的自定义Dense层?

在tf.keras中实现带可训练稀疏权重的自定义Dense层

核心思路

因为你的权重矩阵稀疏模式固定,可以仅对非零位置的权重值进行训练,稀疏张量的indices和dense_shape保持不变,仅更新values部分。这样既节省内存占用,又能通过稀疏矩阵乘法减少计算量。

实现步骤与代码示例

1. 自定义层实现

继承tf.keras.layers.Layer,在初始化阶段传入固定的稀疏模式参数,仅将非零权重值设为可训练变量:

import tensorflow as tf

class SparseDense(tf.keras.layers.Layer):
    def __init__(self, units, sparse_indices, dense_shape, activation=None, **kwargs):
        super().__init__(**kwargs)
        self.units = units
        # 固定的非零权重位置,形状为[非零元素数量, 2]
        self.sparse_indices = sparse_indices
        # 权重矩阵的稠密形状,格式为[输入维度, 输出维度]
        self.dense_shape = dense_shape
        self.activation = tf.keras.activations.get(activation)

    def build(self, input_shape):
        # 仅为非零权重值创建可训练变量
        self.sparse_weights = self.add_weight(
            name="sparse_weights",
            shape=(self.sparse_indices.shape[0],),
            initializer="glorot_uniform",
            trainable=True
        )
        super().build(input_shape)

    def call(self, inputs):
        # 构建稀疏权重张量
        sparse_matrix = tf.sparse.SparseTensor(
            indices=self.sparse_indices,
            values=self.sparse_weights,
            dense_shape=self.dense_shape
        )
        # 执行稀疏-稠密矩阵乘法,替代普通稠密矩阵运算
        outputs = tf.sparse.sparse_dense_matmul(inputs, sparse_matrix)
        if self.activation is not None:
            outputs = self.activation(outputs)
        return outputs

2. 使用示例

假设输入维度为1000,输出维度为200,已知权重矩阵有500个非零值(实际场景替换为你预先确定的稀疏位置):

# 模拟固定稀疏模式(实际用你已知的非零位置)
input_dim = 1000
units = 200
num_nonzero = 500
sparse_indices = tf.random.uniform(shape=(num_nonzero, 2), minval=0, maxval=input_dim, dtype=tf.int64)
# 去重避免重复索引
sparse_indices = tf.unique(tf.reshape(sparse_indices, [-1]))[0]
sparse_indices = tf.reshape(sparse_indices, [-1, 2])
dense_shape = [input_dim, units]

# 初始化层并测试
sparse_dense_layer = SparseDense(units=units, sparse_indices=sparse_indices, dense_shape=dense_shape, activation="relu")
test_input = tf.random.normal(shape=(32, input_dim))  # 批量大小32
output = sparse_dense_layer(test_input)
print(output.shape)  # 输出形状为(32, 200)

3. 训练兼容性说明

  • 层中的sparse_weights是标准可训练变量,TensorFlow自动微分机制会正常计算其梯度,无需额外处理。
  • 稀疏矩阵乘法tf.sparse.sparse_dense_matmul原生支持自动微分,训练流程和普通Dense层完全一致,可直接使用model.compile()、model.fit()等接口。

关键注意事项

  • 稀疏模式必须提前确定,训练过程中不可修改,否则会增加实现复杂度。
  • 若稀疏模式有特定规则(如分块稀疏),可针对性优化存储格式,但通用场景下SparseTensor已能满足需求。
  • 内存节省效果:相比稠密矩阵,稀疏张量仅存储非零值和位置,可节省(输入维度*输出维度 - 非零元素数量)*4字节(按float32计算)。

参考资料(TensorFlow官方文档内容整理)

  • tf.sparse.SparseTensor:用于高效表示稀疏数据,固定形状与非零位置后仅需更新值
  • tf.sparse.sparse_dense_matmul:高效的稀疏-稠密矩阵乘法运算,支持自动微分
  • tf.keras.layers.Layer自定义层规范:通过add_weight方法定义可训练变量

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 11:16:01