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

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"

疑问

  1. 是否有人遇到过同类问题?
  2. 该问题是否需要上报到对应开源仓库?

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.01 04:54:05