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

基于数组输入列表的TensorFlow生成器最优实现方法

利用TensorFlow原生函数优化数据生成器效率

背景

你当前基于TensorFlow/Keras构建的多输入深度学习模型如下:

inps = [] 
features = [] 

for i in range(number_windows):

    inp = Input(shape=(window_length,), name=f"input_{i}")
    inps.append(inp)

    feat = Dense(25)(inp)
    feat = BatchNormalization()(feat)
    feat = LeakyReLU()(feat)
    features.append(feat)

comb = concatenate(features)
comb = Dropout(0.50)(comb)

top = Dense(512)(comb)
top = BatchNormalization()(top)
top = LeakyReLU()(top)
top = Dropout(0.40)(top)

top = Dense(256)(top)

emb = EmbeddingLayer()(top)
    
top = BatchNormalization()(top)
top = LeakyReLU()(top)
top = Dropout(0.25)(top)

classification = Dense(n_classes, activation='softmax', name='classification')(top)

mdl = Model(inputs=inps, outputs=[emb, classification])

其中EmbeddingLayer是自定义L2归一化层。

你已经将批量片段生成优化为逐行生成的版本(注:原代码中spectra_matrix应为data_matrix,已修正):

def data_loading_generator(
    data_matrix: np.typing.NDArray,
    data_labels: np.typing.NDArray,
    window_length,
    dw
):
    num_rows = data_matrix.shape[0]
    for row_number in range(0, num_rows):
        data_segments = segment_data(
            data_matrix[row_number, :], w=window_length, dw=dw
        )
        yield (
            {f"input_{ii}": data_segments[ii, :] for ii in range(data_segments.shape[0])},
            (
                {
                     "embedding_layer": data_labels[row_number],
                     "classification": tf.one_hot(
                         data_labels[row_number], depth=2, dtype=tf.uint16
                     )
                }
            )
        )

基于TensorFlow原生函数的优化方案

1. 用tf.signal.frame替换自定义片段生成

TensorFlow原生的tf.signal.frame可直接在图内完成滑动窗口分割,避免numpy与TensorFlow之间的数据转换开销,替代自定义segment_data函数:

def tf_segment_data(row, window_length, dw):
    # 步长为dw,对应重叠度为window_length - dw
    segments = tf.signal.frame(row, frame_length=window_length, frame_step=dw)
    return segments

2. 切换到tf.data.Dataset替代Python生成器

Python生成器受GIL限制,tf.data.Dataset是TensorFlow原生数据管道,支持并行加载、异步预取、向量化操作,效率提升明显。构建步骤如下:

步骤1:从numpy数据创建基础数据集

# 每条数据对应(行数据, 标签)
dataset = tf.data.Dataset.from_tensor_slices((data_matrix, data_labels))

步骤2:添加预处理映射

预先定义输入键列表,避免动态生成的潜在开销,同时完成片段分割与标签转换:

# 预定义模型输入的键名
input_keys = [f"input_{i}" for i in range(number_windows)]

def preprocess_fn(row, label):
    # 生成窗口片段
    segments = tf.signal.frame(row, frame_length=window_length, frame_step=dw)
    # 确保片段数量与模型输入数量匹配
    segments = tf.ensure_shape(segments, [number_windows, window_length])
    # 转换为模型需要的多输入字典
    input_dict = dict(zip(input_keys, tf.unstack(segments)))
    # 处理输出标签
    output_dict = {
        "embedding_layer": label,
        "classification": tf.one_hot(label, depth=2, dtype=tf.int32)
    }
    return input_dict, output_dict

# 并行执行预处理,AUTOTUNE自动适配CPU资源
dataset = dataset.map(preprocess_fn, num_parallel_calls=tf.data.AUTOTUNE)

步骤3:添加批处理与预取

让GPU训练与CPU数据加载并行,消除数据等待瓶颈:

batch_size = 32
dataset = dataset.batch(batch_size)
# 预取后续批次数据
dataset = dataset.prefetch(tf.data.AUTOTUNE)

3. 内存与数据类型优化

  • 将标签的tf.uint16替换为tf.int32,减少类型转换开销;
  • 提前将data_matrix转换为tf.float32张量,避免图内类型转换;
  • 若数据集过大,可使用tf.data.Dataset.from_generator结合内存映射文件,但优先使用纯TensorFlow操作。

完整优化后的数据管道示例

def build_tf_data_pipeline(data_matrix, data_labels, window_length, dw, number_windows, batch_size=32):
    # 转换为TensorFlow原生张量
    data_matrix = tf.convert_to_tensor(data_matrix, dtype=tf.float32)
    data_labels = tf.convert_to_tensor(data_labels, dtype=tf.int32)
    
    # 创建基础数据集
    dataset = tf.data.Dataset.from_tensor_slices((data_matrix, data_labels))
    
    # 预定义输入键
    input_keys = [f"input_{i}" for i in range(number_windows)]
    
    # 预处理函数
    def preprocess_fn(row, label):
        segments = tf.signal.frame(row, frame_length=window_length, frame_step=dw)
        segments = tf.ensure_shape(segments, [number_windows, window_length])
        input_dict = dict(zip(input_keys, tf.unstack(segments)))
        output_dict = {
            "embedding_layer": label,
            "classification": tf.one_hot(label, depth=2, dtype=tf.int32)
        }
        return input_dict, output_dict
    
    # 构建高效数据管道
    dataset = dataset.map(preprocess_fn, num_parallel_calls=tf.data.AUTOTUNE)
    dataset = dataset.batch(batch_size)
    dataset = dataset.prefetch(tf.data.AUTOTUNE)
    
    return dataset

使用时直接传入模型训练:

train_dataset = build_tf_data_pipeline(data_matrix, data_labels, window_length, dw, number_windows)
mdl.fit(train_dataset, epochs=10)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 10:25:22