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

如何提取SparkNLP WordEmbeddingsModel生成的嵌入以适配Keras/TensorFlow的RNN模型

阿尔巴尼亚语文本分类:SparkNLP嵌入转TensorFlow/Keras RNN输入

我在处理阿尔巴尼亚语(sq)维基数据的文本分类任务,选用SparkNLP的w2v_cc_300d预训练嵌入模型,已通过WordEmbeddingsModel将句子转换为嵌入,但不知道如何将这些嵌入处理为Keras/TensorFlow构建的RNN模型的输入。

我的数据集包含text和label两列,目前已完成以下步骤:

已完成的预处理流程

1. 初始化Spark会话并加载数据

# 启动Spark会话(启用GPU)
spark = sparknlp.start(gpu=True)

# 将训练集转换为Spark DataFrame
spark_train_df = spark.createDataFrame(train)

数据预览:

textlabel
Joy Adowaa Buolam...0
Ajo themeloi "Alg...1
Buolamwini lindi ...1
Kur ishte 9 vjeç,...0
Si një studente u...1

2. 定义SparkNLP预处理Pipeline

# 文档组装器
document = DocumentAssembler()\
    .setInputCol("text")\
    .setOutputCol("document")

# 分词器
tokenizer = Tokenizer() \
    .setInputCols(["document"]) \
    .setOutputCol("token")

# 预训练词嵌入模型
embeddings = WordEmbeddingsModel\
    .pretrained("w2v_cc_300d", "sq")\
    .setInputCols(["document", "token"])\
    .setOutputCol("embeddings")

# 构建Pipeline
pipeline = Pipeline(stages=[document, tokenizer, embeddings])

# 拟合并转换数据
model = pipeline.fit(spark_train_df)
result = model.transform(spark_train_df)

转换后的数据预览:

textlabeldocumenttokenembeddings
Joy Adowaa Buolam...0[{document, 0, 13...[{token, 0, 2, Jo...[{word_embeddings...
Ajo themeloi "Alg...1[{document, 0, 13...[{token, 0, 2, Aj...[{word_embeddings...
Buolamwini lindi ...1[{document, 0, 94...[{token, 0, 9, Bu...[{word_embeddings...
Kur ishte 9 vjeç,...0[{document, 0, 12...[{token, 0, 2, Ku...[{word_embeddings...
Si një studente u...1[{document, 0, 15...[{token, 0, 1, Si...[{word_embeddings...

3. 输出数据Schema

result.printSchema()

输出:

root
 |-- text: string (nullable = true)
 |-- label: long (nullable = true)
 |-- document: array (nullable = true)
 |    |-- element: struct (containsNull = true)
 |    |    |-- annotatorType: string (nullable = true)
 |    |    |-- begin: integer (nullable = false)
 |    |    |-- end: integer (nullable = false)
 |    |    |-- result: string (nullable = true)
 |    |    |-- metadata: map (nullable = true)
 |    |    |    |-- key: string
 |    |    |    |-- value: string (valueContainsNull = true)
 |    |    |-- embeddings: array (nullable = true)
 |    |    |    |-- element: float (containsNull = false)
 |-- token: array (nullable = true)
 |    |-- element: struct (containsNull = true)
 |    |    |-- annotatorType: string (nullable = true)
 |    |    |-- begin: integer (nullable = false)
 |    |    |-- end: integer (nullable = false)
 |    |    |-- result: string (nullable = true)
 |    |    |-- metadata: map (nullable = true)
 |    |    |    |-- key: string
 |    |    |    |-- value: string (valueContainsNull = true)
 |    |    |-- embeddings: array (nullable = true)
 |    |    |    |-- element: float (containsNull = false)
 |-- embeddings: array (nullable = true)
 |    |-- element: struct (containsNull = true)
 |    |    |-- annotatorType: string (nullable = true)
 |    |    |-- begin: integer (nullable = false)
 |    |    |-- end: integer (nullable = false)
 |    |    |-- result: string (nullable = true)
 |    |    |-- metadata: map (nullable = true)
 |    |    |    |-- key: string
 |    |    |    |-- value: string (valueContainsNull = true)
 |    |    |-- embeddings: array (nullable = true)
 |    |    |    |-- element: float (containsNull = false)

embeddings列的数据类型:

result.schema["embeddings"].dataType

输出:

ArrayType(StructType([StructField('annotatorType', StringType(), True), StructField('begin', IntegerType(), False), StructField('end', IntegerType(), False), StructField('result', StringType(), True), StructField('metadata', MapType(StringType(), StringType(), True), True), StructField('embeddings', ArrayType(FloatType(), False), True)]), True)

嵌入转换为RNN输入的解决方案

1. 提取纯嵌入向量矩阵

SparkNLP输出的embeddings列是包含元数据的结构体数组,我们需要提取每个token对应的300维嵌入向量:

from pyspark.sql.functions import udf, col
from pyspark.sql.types import ArrayType, FloatType

# 定义UDF提取每个token的嵌入向量
extract_embeddings = udf(lambda x: [token['embeddings'] for token in x], ArrayType(ArrayType(FloatType())))

# 生成仅包含嵌入和标签的DataFrame
processed_df = result.select(
    extract_embeddings(col("embeddings")).alias("token_embeddings"),
    col("label")
)

2. 统一序列长度

RNN模型要求输入序列长度一致,需对嵌入序列进行截断或填充:

from pyspark.sql.functions import array, lit, slice, size

# 设定最大序列长度(根据数据集统计调整)
MAX_SEQ_LENGTH = 100

# 填充零向量到固定长度,过长则截断
processed_df = processed_df.withColumn(
    "padded_embeddings",
    slice(
        array(*[lit([0.0]*300) for _ in range(MAX_SEQ_LENGTH)]),
        1,
        MAX_SEQ_LENGTH - size(col("token_embeddings"))
    ).concat(col("token_embeddings"))
).withColumn(
    "padded_embeddings",
    slice(col("padded_embeddings"), 1, MAX_SEQ_LENGTH)
)

3. 转换为TensorFlow兼容格式

将Spark DataFrame转换为NumPy数组或TensorFlow Dataset:

import numpy as np
import tensorflow as tf

# 收集数据并转换为NumPy数组
data = processed_df.select("padded_embeddings", "label").collect()
X = np.array([row['padded_embeddings'] for row in data], dtype=np.float32)
y = np.array([row['label'] for row in data], dtype=np.int32)

# 转换为TensorFlow Dataset(推荐,适合大规模数据)
dataset = tf.data.Dataset.from_tensor_slices((X, y)).shuffle(1000).batch(32)

4. 构建并训练RNN模型

使用处理好的输入构建分类模型:

from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import LSTM, Dense

model = Sequential([
    # 无需额外Embedding层,直接使用预训练嵌入
    LSTM(64, input_shape=(MAX_SEQ_LENGTH, 300)),
    Dense(32, activation='relu'),
    Dense(1, activation='sigmoid')  # 二分类任务,多分类请调整输出维度和激活函数
])

model.compile(optimizer='adam', loss='binary_crossentropy', metrics=['accuracy'])
model.fit(dataset, epochs=10)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 05:31:20