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

在Dask分区并行处理中应用TextVectorization和StringLookup遇序列化错误

问题描述

我正尝试构建管道,对无法载入内存的超大规模数据集并行写入TFRecord文件。此前已多次用Dask完成这类任务,但新数据集需在模型外应用TextVectorization和StringLookup(模型内运行会导致GPU性能瓶颈)。目标是在单分区串行应用这些工具后,用Dask并行处理各分区。尝试多种写法及tf.keras序列化反序列化方法均未成功,出现序列化错误。

当前代码

from dask.distributed import LocalCluster, Client
import logging
import dask
import json
import os
import pickle
import tensorflow
from tensorflow.keras.layers.experimental.preprocessing import StringLookup, TextVectorization

cluster = LocalCluster()

client = Client(cluster)

# Setup logging
logging.basicConfig(format='%(levelname)s:%(message)s', level=logging.INFO)
 
def load_and_broadcast_vectorizers(model_dir):
    logging.info("Loading and broadcasting vectorizers...")
    path = os.path.join(model_dir, 'vectorizers_saved')
    vectorizers = {}
   
    for file in os.listdir(path):
        if file.endswith("_vectorizer_config.json"):
            key = file[:-23]
            with open(os.path.join(path, f'{key}_vectorizer_config.json'), 'r') as f:
                config = json.load(f)
           
            vectorizer_class = config['vectorizer_class']
            config.pop('vectorizer_class', None)  # Remove the 'vectorizer_class' key
 
            if vectorizer_class == 'StringLookup':
                vectorizer = StringLookup.from_config(config)
            elif vectorizer_class == 'TextVectorization':
                vectorizer = TextVectorization.from_config(config)
            else:
                raise ValueError(f"Unknown vectorizer class: {vectorizer_class}")
 
            vectorizers[key] = vectorizer
    logging.info("Vectorizers loaded and broadcasted successfully.")
    return client.scatter(vectorizers, broadcast=True)
 
def apply_vectorizers_to_partition(partition, vectorizers):
    logging.info("Applying vectorizers to partition...")
    for vectorizer_name, vectorizer in vectorizers.items():
        partition[vectorizer_name] = vectorizer.transform(partition[vectorizer_name])
    return partition
 
def apply_vectorizers_in_parallel(vectorizers, df):
    logging.info("Applying vectorizers in parallel...")
    return df.map_partitions(apply_vectorizers_to_partition, vectorizers)
 
def write_partition_to_tfrecord(partition, output_path, partition_label, partition_id):
    logging.info(f"Writing {partition_label} partition {partition_id} to TFRecord...")
    file_name_tfrecord = f'{partition_label}_partition_{partition_id}.tfrecord'
    output_file_path_tfrecord = os.path.join(output_path, file_name_tfrecord)
    with tf.io.TFRecordWriter(output_file_path_tfrecord) as writer:
        for row in partition.itertuples():
                        # Extract features and label from each row
            features = {
                'input_1': tf.train.Feature(int64_list=tf.train.Int64List(value=[row['input_1']])),
                'input_2': tf.train.Feature(float_list=tf.train.FloatList(value=[row['input_2']])),
                'input_3': tf.train.Feature(int64_list=tf.train.Int64List(value=row['input_3'])),
                'input_4': tf.train.Feature(int64_list=tf.train.Int64List(value=[row['input_4']])),
                'input_5': tf.train.Feature(int64_list=tf.train.Int64List(value=[row['input_5']]))
    }

            label = tf.train.Feature(float_list=tf.train.FloatList(value=[row['label_col']]))
            example = tf.train.Example(features=tf.train.Features(feature={**features, **{'label': label}}))
           writer.write(example.SerializeToString())
    logging.info(f"{partition_label} partition {partition_id} written to TFRecord.")
 
def write_partition_to_parquet(partition, output_path, partition_label, partition_id):
    logging.info(f"Writing {partition_label} partition {partition_id} to Parquet...")
    selected_columns = partition[['personuuid', 'claimnum', 'claimrowid']]
    file_name_parquet = f'{partition_label}_partition_{partition_id}.parquet.snappy'
    output_file_path_parquet = os.path.join(output_path, file_name_parquet)
    selected_columns.to_parquet(output_file_path_parquet, compression='snappy')
    logging.info(f"{partition_label} partition {partition_id} written to Parquet.")
 
def write_vectorized_partitions_to_files(vectorizers, df, output_path, partition_label):
    logging.info(f"Writing {partition_label} vectorized partitions to files...")
    dask_tasks = []
    for i, partition in enumerate(df.to_delayed()):
        tfrecord_task = dask.delayed(write_partition_to_tfrecord)(partition, output_path, partition_label, i)
        parquet_task = dask.delayed(write_partition_to_parquet)(partition, output_path, partition_label, i)
        dask_tasks.extend([tfrecord_task, parquet_task])
    dask.compute(*dask_tasks)
    logging.info(f"{partition_label} vectorized partitions written to files successfully.")
 
def process_data(model_dir, df, output_path):
    logging.info("Processing data...")
    vectorizers = load_and_broadcast_vectorizers(model_dir)
   
    if vectorizers is None:
        logging.error("Data processing failed due to missing vectorizers.")
        return
 
    train_df, test_df = df.random_split([0.8, 0.2], random_state=42)
    train_df_vectorized = apply_vectorizers_in_parallel(vectorizers, train_df)
    test_df_vectorized = apply_vectorizers_in_parallel(vectorizers, test_df)
    write_vectorized_partitions_to_files(vectorizers, train_df_vectorized, os.path.join(output_path, 'train'), 'train')
    write_vectorized_partitions_to_files(vectorizers, test_df_vectorized, os.path.join(output_path, 'test'), 'test')
    logging.info("Data processed successfully.")

错误信息

TypeError: ('Could not serialize object of type StringLookup', '<keras.layers.preprocessing.string_lookup.StringLookup object at 0x7f7f1d79e9a0>')
解决方案

问题根源是Dask无法序列化Keras的StringLookup/TextVectorization层实例,直接广播这些对象会触发序列化失败。解决思路是让每个Dask Worker自行加载vectorizer配置并重建实例,而不是广播已经创建好的对象。

修改步骤

  1. 重构配置获取函数:只返回vectorizer配置文件的路径,不创建实例
def get_vectorizer_configs_path(model_dir):
    logging.info("Getting vectorizer configs path...")
    path = os.path.join(model_dir, 'vectorizers_saved')
    if not os.path.exists(path):
        raise FileNotFoundError(f"Vectorizer configs path not found: {path}")
    return path
  1. 修改分区处理函数:让每个Worker在处理分区时自行加载并重建vectorizer
def apply_vectorizers_to_partition(partition, vectorizer_configs_path):
    logging.info("Applying vectorizers to partition...")
    vectorizers = {}
    # 加载并重建vectorizer
    for file in os.listdir(vectorizer_configs_path):
        if file.endswith("_vectorizer_config.json"):
            key = file[:-23]
            with open(os.path.join(vectorizer_configs_path, f'{key}_vectorizer_config.json'), 'r') as f:
                config = json.load(f)
           
            vectorizer_class = config['vectorizer_class']
            config.pop('vectorizer_class', None)
 
            if vectorizer_class == 'StringLookup':
                vectorizer = StringLookup.from_config(config)
            elif vectorizer_class == 'TextVectorization':
                vectorizer = TextVectorization.from_config(config)
            else:
                raise ValueError(f"Unknown vectorizer class: {vectorizer_class}")
 
            vectorizers[key] = vectorizer
    # 应用vectorizer到分区
    for vectorizer_name, vectorizer in vectorizers.items():
        partition[vectorizer_name] = vectorizer.transform(partition[vectorizer_name])
    return partition
  1. 调整并行应用函数:传入配置路径而非vectorizer实例
def apply_vectorizers_in_parallel(vectorizer_configs_path, df):
    logging.info("Applying vectorizers in parallel...")
    return df.map_partitions(apply_vectorizers_to_partition, vectorizer_configs_path)
  1. 更新主处理函数:移除vectorizer实例的传递逻辑
def process_data(model_dir, df, output_path):
    logging.info("Processing data...")
    vectorizer_configs_path = get_vectorizer_configs_path(model_dir)
   
    if vectorizer_configs_path is None:
        logging.error("Data processing failed due to missing vectorizer config path.")
        return
 
    train_df, test_df = df.random_split([0.8, 0.2], random_state=42)
    train_df_vectorized = apply_vectorizers_in_parallel(vectorizer_configs_path, train_df)
    test_df_vectorized = apply_vectorizers_in_parallel(vectorizer_configs_path, test_df)
    write_vectorized_partitions_to_files(train_df_vectorized, os.path.join(output_path, 'train'), 'train')
    write_vectorized_partitions_to_files(test_df_vectorized, os.path.join(output_path, 'test'), 'test')
    logging.info("Data processed successfully.")
  1. 更新文件写入函数:移除多余的vectorizer参数
def write_vectorized_partitions_to_files(df, output_path, partition_label):
    logging.info(f"Writing {partition_label} vectorized partitions to files...")
    dask_tasks = []
    for i, partition in enumerate(df.to_delayed()):
        tfrecord_task = dask.delayed(write_partition_to_tfrecord)(partition, output_path, partition_label, i)
        parquet_task = dask.delayed(write_partition_to_parquet)(partition, output_path, partition_label, i)
        dask_tasks.extend([tfrecord_task, parquet_task])
    dask.compute(*dask_tasks)
    logging.info(f"{partition_label} vectorized partitions written to files successfully.")

额外注意事项

  • 确保所有Dask Worker都能访问到vectorizer配置文件所在路径(分布式集群需使用共享存储,如NFS、S3等)。
  • 如果vectorizer依赖额外文件(如TextVectorization的词汇表),需确保这些文件也在Worker可访问路径下,并在加载配置时正确关联。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.19 23:57:05