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

TensorFlow自定义Keras层训练时动态添加变量/权重实现可扩展嵌入层

训练过程中向自定义Keras层动态添加变量/权重的实现方案

问题描述

能否在训练过程中给自定义Keras层添加新的变量/权重?我需要这个功能来构建可扩展的嵌入层。

尝试过的方法及问题

我试过用tf.py_function,但它没法追踪新增的权重,会抛出无梯度可用的异常。用tf.keras.layers.Lambda层也有类似问题,触发如下警告:

WARNING:tensorflow:
The following Variables were used a Lambda layer's call (lambda), but
are not present in its tracked objects:
<tf.Variable 'model/boundless_embedding/block_0:0' shape=(2, 3) dtype=float32, numpy=
array([[-0.4273603 , -0.21527308, -0.3599682 ],
[-0.04728699, 0.71099687, -0.04171979]], dtype=float32)>
<tf.Variable 'model/boundless_embedding/block_1:0' shape=(2, 3) dtype=float32, numpy=
array([[ 0.44539702, 0.3135407 , -0.94582325],
[ 0.42753923, -0.15626878, 0.5873704 ]], dtype=float32)>
It is possible that this is intended behavior, but it is more likely
an omission. This is a strong indication that this layer should be
formulated as a subclassed Layer rather than a Lambda layer.

原始代码

import tensorflow as tf
import numpy


class BoundlessEmbedding(tf.keras.layers.Layer):
    def __init__(self, dimension, block_size=2**20):
        super().__init__(dynamic=True)

        self._block_size = block_size
        self._dimension = dimension
        self.block_weights = []
        self.lookup = tf.keras.layers.Lambda(lambda x: self._lookup(x))

    def call(self, x, training=None, mask=None):
        if training:
            with tf.init_scope():
                self._maybe_expand(x)

        def lookup(x_):
            return self._lookup(x_)

        # y = tf.py_function(lookup, [x], tf.float32)
        y = self.lookup(x)

        return tf.reduce_sum(y * y, axis=1)

    def compute_output_shape(self, input_shape):
        return input_shape + [self._dimension]

    def _maybe_expand(self, x):
        maximum = tf.math.reduce_max(x)
        while maximum >= len(self.block_weights) * self._block_size:
            id_ = len(self.block_weights)
            weight = self.add_weight(
                    name=f'block_{id_}', dtype=tf.float32,
                    shape=[self._block_size, self._dimension])
            self.block_weights.append(weight)

    def _lookup(self, x):
        # TODO Remove this reshape
        x = tf.reshape(x, [-1])

        valid = x < len(self.block_weights) * self._block_size
        x = tf.where(valid, x, 0)

        y = tf.zeros([x.shape[0], self._dimension], dtype=tf.float32)
        block_ids = x // self._block_size
        block_offsets = x % self._block_size

        for i, weight in enumerate(self.block_weights):
            idx = tf.where(tf.math.equal(block_ids, i))
            offsets = tf.gather(block_offsets, tf.reshape(idx, [-1]))
            values = tf.gather(weight, offsets)
            y = tf.tensor_scatter_nd_update(y, idx, values)
        return tf.reshape(y, [-1, self._dimension])


def build_dataset():
    x = numpy.array([[0, 1, 2, 3], [4, 5, 6, 7]], dtype=numpy.int32)
    y = numpy.array([[0, 0, 0, 0], [0, 0, 0, 0]], dtype=numpy.float32)
    x = tf.data.Dataset.from_tensor_slices(x)
    y = tf.data.Dataset.from_tensor_slices(y)
    return tf.data.Dataset.zip((x, y))


def main():
    #tf.enable_eager_execution()

    embedding = BoundlessEmbedding(3, 2)
    x = tf.keras.Input(name="x", shape=[None], dtype=tf.int32)
    y = embedding(x)

    model = tf.keras.Model(inputs=x, outputs=y)
    model.compile(optimizer='sgd', loss='mse')

    dataset = build_dataset()
    model.fit(dataset)


if __name__ == '__main__':
    main()

问题根源与解决思路

Lambda层和tf.py_function无法正确追踪动态添加的权重,因为它们不属于Layer的可追踪对象集合。自定义子类化Layer时,需要确保新增权重被Keras的变量追踪机制正确捕获,同时避免在计算图构建阶段引入动态操作冲突。

修正后的代码实现

核心改动点:

  • 移除Lambda层,直接在call方法中调用_lookup,避免权重追踪丢失
  • 调整_maybe_expand的执行时机,确保在Eager模式或Graph模式下都能正确初始化权重
  • 优化张量操作的可读性
import tensorflow as tf
import numpy as np


class BoundlessEmbedding(tf.keras.layers.Layer):
    def __init__(self, dimension, block_size=2**20):
        super().__init__()
        self._block_size = block_size
        self._dimension = dimension
        # 使用列表存储权重块,Keras会自动追踪列表内的Variable
        self.block_weights = []

    def call(self, x, training=None):
        if training:
            # 在训练模式下检查并扩展权重块
            self._maybe_expand(x)
        
        # 直接调用_lookup,无需Lambda或py_function
        embeddings = self._lookup(x)
        # 保持原代码的输出逻辑:embedding的平方和
        return tf.reduce_sum(embeddings * embeddings, axis=-1)

    def _maybe_expand(self, x):
        # 计算输入中的最大索引,确定是否需要扩展权重块
        max_idx = tf.math.reduce_max(x)
        current_max_id = len(self.block_weights) * self._block_size
        
        # 循环扩展直到覆盖最大索引
        while max_idx >= current_max_id:
            block_id = len(self.block_weights)
            # 添加可训练权重,Keras会自动追踪
            new_block = self.add_weight(
                name=f'block_{block_id}',
                shape=(self._block_size, self._dimension),
                dtype=tf.float32,
                trainable=True
            )
            self.block_weights.append(new_block)
            current_max_id += self._block_size

    def _lookup(self, x):
        # 保留原逻辑,处理任意形状的输入
        original_shape = tf.shape(x)
        x_flat = tf.reshape(x, [-1])
        
        # 过滤超出范围的索引,替换为0(可根据需求调整)
        valid_mask = x_flat < len(self.block_weights) * self._block_size
        x_flat = tf.where(valid_mask, x_flat, 0)
        
        # 计算每个索引所属的块ID和偏移量
        block_ids = x_flat // self._block_size
        offsets = x_flat % self._block_size
        
        # 初始化输出张量
        embeddings_flat = tf.zeros((tf.shape(x_flat)[0], self._dimension), dtype=tf.float32)
        
        # 遍历权重块,填充对应索引的嵌入向量
        for block_idx, weight_block in enumerate(self.block_weights):
            # 找到属于当前块的索引位置
            mask = tf.equal(block_ids, block_idx)
            idx_positions = tf.where(mask)
            idx_flat = tf.reshape(idx_positions, [-1])
            
            # 获取对应偏移量的嵌入向量
            selected_offsets = tf.gather(offsets, idx_flat)
            selected_embeddings = tf.gather(weight_block, selected_offsets)
            
            # 更新输出张量
            embeddings_flat = tf.tensor_scatter_nd_update(
                embeddings_flat, idx_positions, selected_embeddings
            )
        
        # 恢复原始形状
        return tf.reshape(embeddings_flat, tf.concat([original_shape, [self._dimension]], axis=0))


def build_dataset():
    x = np.array([[0, 1, 2, 3], [4, 5, 6, 7]], dtype=np.int32)
    y = np.array([[0, 0, 0, 0], [0, 0, 0, 0]], dtype=np.float32)
    x_ds = tf.data.Dataset.from_tensor_slices(x)
    y_ds = tf.data.Dataset.from_tensor_slices(y)
    return tf.data.Dataset.zip((x_ds, y_ds))


def main():
    embedding = BoundlessEmbedding(3, 2)
    x_input = tf.keras.Input(name="x", shape=[None], dtype=tf.int32)
    output = embedding(x_input)
    
    model = tf.keras.Model(inputs=x_input, outputs=output)
    model.compile(optimizer='sgd', loss='mse')
    
    dataset = build_dataset()
    model.fit(dataset, epochs=1)


if __name__ == '__main__':
    main()

关键说明

  1. 权重追踪:直接在子类化Layer中使用add_weight添加的变量会被Keras自动追踪,无需额外处理。Lambda层无法感知Layer内部的变量,因此会触发警告。
  2. 动态扩展逻辑:_maybe_expand在训练模式下执行,确保遇到新的索引时自动扩展权重块。tf.init_scope()可以移除,因为add_weight在训练时调用会自动处理变量初始化。
  3. 计算图兼容性:直接在call中调用_lookup避免了tf.py_function的梯度丢失问题,所有操作都在TensorFlow计算图中,支持自动微分。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.14 21:56:03