在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配置并重建实例,而不是广播已经创建好的对象。
修改步骤
- 重构配置获取函数:只返回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
- 修改分区处理函数:让每个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
- 调整并行应用函数:传入配置路径而非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)
- 更新主处理函数:移除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.")
- 更新文件写入函数:移除多余的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
相关产品推荐
相关产品推荐

