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

如何在Tensorflow中构建段落序列RNN分类模型并学习段落嵌入

当然可以实现,核心就是把「段落」当成新的「词汇」来处理

这个思路完全可行,本质和Tensorflow官方RNN文本分类教程的逻辑一致,只是把输入的基本单元从「词」换成了「段落」,具体实现步骤如下:

1. 先给段落做索引编码(类比构建词表)

把语料里所有出现过的段落收集起来,给每个唯一段落分配一个专属索引,就像处理词时构建词表一样。用Tensorflow的StringLookup层就能快速实现:

import tensorflow as tf
import numpy as np

# 示例文档数据:每个元素是(段落列表, 分类标签)
docs = [
    (["特定领域段落A内容", "特定领域段落B内容"], 0),
    (["特定领域段落C内容", "特定领域段落D内容", "特定领域段落E内容"], 1),
    # 替换成你的真实语料
]

# 提取所有段落,构建段落集合
all_paragraphs = []
for paras, _ in docs:
    all_paragraphs.extend(paras)

# 构建段落到索引的映射层
paragraph_lookup = tf.keras.layers.StringLookup(
    vocabulary=list(set(all_paragraphs)), 
    mask_token=None,
    output_mode="int"
)

2. 处理文档输入序列(类比处理句子的词序列)

把每个文档的段落列表转换成索引序列,然后统一长度(截断过长的文档,补齐过短的文档):

max_paragraphs_per_doc = 50  # 根据你的文档情况调整,比如取前50个段落

X = []
y = []
for paras, label in docs:
    # 把段落转成索引
    para_indices = paragraph_lookup(paras).numpy()
    # 统一序列长度
    padded_indices = tf.keras.preprocessing.sequence.pad_sequences(
        [para_indices], 
        maxlen=max_paragraphs_per_doc,
        padding="post",
        truncating="post"
    )[0]
    X.append(padded_indices)
    y.append(label)

# 转成模型能接受的数组格式
X = np.array(X)
y = np.array(y)

3. 构建带段落嵌入的RNN分类模型(替换词嵌入为段落嵌入)

直接用Embedding层学习段落向量,后续接RNN层和分类层,和官方教程的结构几乎完全一样:

vocab_size = len(paragraph_lookup.get_vocabulary())
embedding_dim = 128  # 段落嵌入的维度,可根据需求调整
rnn_units = 64  # RNN单元数,可调整

model = tf.keras.Sequential([
    # 段落嵌入层:和词嵌入层逻辑完全一致,训练时自动学习段落向量
    tf.keras.layers.Embedding(
        input_dim=vocab_size,
        output_dim=embedding_dim,
        mask_zero=True  # 让RNN忽略补齐的0向量
    ),
    # RNN层:用LSTM/GRU都可以,根据需求选择
    tf.keras.layers.LSTM(rnn_units),
    # 分类输出层:二分类用sigmoid,多分类用softmax
    tf.keras.layers.Dense(1, activation="sigmoid")
])

# 编译并训练模型
model.compile(optimizer="adam", loss="binary_crossentropy", metrics=["accuracy"])
model.fit(X, y, epochs=10, validation_split=0.2)

一些实用注意事项

  • 如果你的语料中段落总数极大,嵌入层的参数会很多,可以给Embedding层添加kernel_regularizer=tf.keras.regularizers.l2(0.01)来防止过拟合;或者用Hashing层结合Embedding做哈希嵌入(但会有冲突,特定领域优先用索引方式)。
  • 如果段落本身很长,也可以先给每个段落单独做编码(比如用小CNN/Transformer把段落转成固定向量),再把这些向量序列输入RNN,但这就不属于直接学习段落嵌入的范畴了,你的需求里直接把段落当单元是完全没问题的。
  • 序列长度max_paragraphs_per_doc要根据你的文档实际段落数调整,比如统计所有文档的段落数分布,取95%分位数作为截断长度。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 14:37:19