AI平台ParameterServerStrategy配合BigQuery连接器报Op未注册错误
问题描述
我正在使用TensorFlow的ParameterServerStrategy对模型训练步骤做并行化处理,基于GCP AI Platform创建集群并启动任务。由于数据集体量极大,我采用了tensorflow-io内置的BigQuery TensorFlow连接器读取数据。
我的脚本参考TensorFlow BigQuery读取器官方文档和TensorFlow ParameterServerStrategy官方文档编写。
该脚本在本地运行完全正常,但部署到AI Platform运行时抛出如下错误:
{"created":"@1633444428.903993309","description":"Error received from peer ipv4:10.46.92.135:2222","file":"external/com_github_grpc_grpc/src/core/lib/surface/call.cc","file_line":1056,"grpc_message":"Op type not registered 'IO>BigQueryClient' in binary running on gke-cml-1005-141531--n1-standard-16-2-644bc3f8-7h8p. Make sure the Op and Kernel are registered in the binary running in this process. Note that if you are loading a saved graph which used ops from tf.contrib, accessing (e.g.) `tf.contrib.resampler` should be done before importing the graph, as contrib ops are lazily registered when the module is first accessed.","grpc_status":5}
已验证信息
- 脚本使用模拟数据在AI Platform上可正常运行
- 本地调用BigQuery连接器可正常运行
- 最初怀疑包含BigQuery连接器的模型编译后,在其他节点调用时触发该问题,但暂未找到修复方案
- 核对所有节点依赖版本完全一致:
- tensorflow : 2.5.0
- tensorflow-io : 0.19.1
- BigQuery连接器在AI Platform的
MirroredStrategy策略下运行完全正常,仅切换为ParameterServerStrategy时出现该问题
最小复现代码
import os from tensorflow_io.bigquery import BigQueryClient from tensorflow_io.bigquery import BigQueryReadSession import tensorflow as tf import multiprocessing import portpicker from tensorflow.keras.layers.experimental import preprocessing from google.cloud import bigquery from tensorflow.python.framework import dtypes import numpy as np import pandas as pd client = bigquery.Client() PROJECT_ID = <your_project> DATASET_ID = 'tmp' TABLE_ID = 'bq_tf_io' BATCH_SIZE = 32 # Bigquery requirements def init_bq_table(): table = '%s.%s.%s' %(PROJECT_ID, DATASET_ID, TABLE_ID) # Create toy_data def create_toy_data(N): x = np.random.random(size = N) y = 0.2 + x + np.random.normal(loc=0, scale = 0.3, size = N) return x, y x, y =create_toy_data(1000) df = pd.DataFrame(data = {'x': x, 'y': y}) job_config = bigquery.LoadJobConfig(write_disposition="WRITE_TRUNCATE",) job = client.load_table_from_dataframe( df, table, job_config=job_config ) job.result() # Create initial data #init_bq_table() CSV_SCHEMA = [ bigquery.SchemaField("x", "FLOAT64"), bigquery.SchemaField("y", "FLOAT64"), ] def transform_row(row_dict): # Trim all string tensors dataset_x = row_dict dataset_x['constant'] = tf.cast(1, tf.float64) # Extract feature column dataset_y = dataset_x.pop('y') #Export as tensor dataset_x = tf.stack([dataset_x[column] for column in dataset_x], axis=-1) return (dataset_x, dataset_y) def read_bigquery(table_name): tensorflow_io_bigquery_client = BigQueryClient() read_session = tensorflow_io_bigquery_client.read_session( "projects/" + PROJECT_ID, PROJECT_ID, TABLE_ID, DATASET_ID, list(field.name for field in CSV_SCHEMA), list(dtypes.double if field.field_type == 'FLOAT64' else dtypes.string for field in CSV_SCHEMA), requested_streams=2) dataset = read_session.parallel_read_rows() return dataset def get_data(): dataset = read_bigquery(TABLE_ID) dataset = dataset.map(transform_row, num_parallel_calls=4) dataset = dataset.batch(BATCH_SIZE).prefetch(2) return dataset cluster_resolver = tf.distribute.cluster_resolver.TFConfigClusterResolver() # parameter server and worker just wait jobs from the coordinator (chief) if cluster_resolver.task_type in ("worker"): worker_config = tf.compat.v1.ConfigProto() server = tf.distribute.Server( cluster_resolver.cluster_spec(), job_name=cluster_resolver.task_type, task_index=cluster_resolver.task_id, config=worker_config, protocol="grpc") server.join() elif cluster_resolver.task_type in ("ps"): server = tf.distribute.Server( cluster_resolver.cluster_spec(), job_name=cluster_resolver.task_type, task_index=cluster_resolver.task_id, protocol="grpc") server.join() elif cluster_resolver.task_type == 'chief': strategy = tf.distribute.experimental.ParameterServerStrategy(cluster_resolver=cluster_resolver) if cluster_resolver.task_type == 'chief': learning_rate = 0.01 with strategy.scope(): # model model_input = tf.keras.layers.Input( shape=(2,), dtype=tf.float64) layer_1 = tf.keras.layers.Dense( 8, activation='relu')(model_input) dense_output = tf.keras.layers.Dense(1)(layer_1) model = tf.keras.Model(model_input, dense_output) #optimizer optimizer=tf.keras.optimizers.SGD(learning_rate=learning_rate) accuracy = tf.keras.metrics.MeanSquaredError() @tf.function def distributed_train_step(iterator): def train_step(x_batch_train, y_batch_train): with tf.GradientTape() as tape: y_predict = model(x_batch_train, training=True) loss_value = tf.keras.losses.MeanSquaredError(reduction=tf.keras.losses.Reduction.NONE)(y_batch_train, y_predict) grads = tape.gradient(loss_value, model.trainable_weights) optimizer.apply_gradients(zip(grads, model.trainable_weights)) accuracy.update_state(y_batch_train, y_predict) return loss_value x_batch_train, y_batch_train = next(iterator) return strategy.run(train_step, args=(x_batch_train, y_batch_train)) coordinator = tf.distribute.experimental.coordinator.ClusterCoordinator(strategy) #test def dataset_fn(_): def create_toy_data(N): x = np.random.random(size = N) y = 0.2 + x + np.random.normal(loc=0, scale = 0.3, size = N) return np.c_[x,y] def toy_transform_row(row): dataset_x = tf.stack([row[0], tf.cast(1, tf.float64)], axis=-1) dataset_y = row[1] return dataset_x, dataset_y N = 1000 data =create_toy_data(N) dataset = tf.data.Dataset.from_tensor_slices(data) dataset = dataset.map(toy_transform_row, num_parallel_calls=4) dataset = dataset.batch(BATCH_SIZE) dataset = dataset.prefetch(2) return dataset @tf.function def per_worker_dataset_fn(): return strategy.distribute_datasets_from_function(lambda x : get_data()) # <-- 切换为该逻辑在AI Platform上报错 #return strategy.distribute_datasets_from_function(dataset_fn) # <-- 切换为该逻辑在AI Platform上正常运行 per_worker_dataset = coordinator.create_per_worker_dataset(per_worker_dataset_fn) # Train model for epoch in range(5): per_worker_iterator = iter(per_worker_dataset) accuracy.reset_states() for step in range(5): coordinator.schedule(distributed_train_step, args=(per_worker_iterator,)) coordinator.join() print ("Finished epoch %d, accuracy is %f." % (epoch, accuracy.result().numpy()))
当在per_worker_dataset_fn()中使用BigQuery连接器生成数据集时触发错误,使用实时生成的模拟数据集则运行正常。
AI Platform集群配置
- runtimeVersion: "2.5"
- pythonVersion: "3.7"
疑问
- 是否有人遇到过同类问题?
- 该问题是否需要上报到对应开源仓库?
内容的提问来源于stack exchange,提问作者Harold G
相关产品推荐
相关产品推荐

