在Databricks中创建Petastorm的make_batch_reader对象报错的解决方法
解决Databricks中Petastorm make_batch_reader的TypeError: missing 'instance' and 'token'参数问题
问题背景
本地环境使用Petastorm的make_batch_reader读取Parquet文件训练模型一切正常,但迁移到Databricks后,使用dbfs://路径调用该方法时触发错误:
TypeError: init() missing 2 required positional arguments: 'instance' and 'token'
错误原因
该错误源于Petastorm在Databricks环境下访问DBFS背后的云存储(如Azure ADLS、AWS S3)时,缺少必要的身份认证参数。本地文件系统访问无需额外认证,但Databricks中的存储访问依赖云服务商的认证机制,默认的dbfs://路径调用未传递所需认证信息,导致底层存储客户端初始化失败。
解决方案
1. Azure Databricks + ADLS Gen2 适配示例
通过Databricks的TokenLibrary获取存储认证令牌,传入storage_options参数:
# 初始化TokenLibrary TokenLibrary = spark._jvm.com.microsoft.azure.synapse.tokenlibrary.TokenLibrary # 替换为你的Linked Service名称 linked_service_name = "your-linked-service-name" # 获取SAS Token sas_token = TokenLibrary.getConnectionString(linked_service_name) # 使用ADLS原生路径(推荐)或dbfs路径 parquet_path = "abfs://container-name@storage-account-name.dfs.core.windows.net/output/scaled.parquet" with make_batch_reader( parquet_path, num_epochs=4, shuffle_row_groups=False, storage_options={"sas_token": sas_token} ) as train_reader: train_ds = make_petastorm_dataset(train_reader).unbatch().map(lambda x: (tf.convert_to_tensor(x))).batch(2) for ele in train_ds: tensor = tf.reshape(ele,(2,1,15)) model.fit(tensor,tensor)
2. AWS Databricks 适配示例
传入S3存储的访问密钥作为认证参数:
# 从Spark配置中读取S3认证信息 storage_options = { "key": spark.conf.get("spark.hadoop.fs.s3a.access.key"), "secret": spark.conf.get("spark.hadoop.fs.s3a.secret.key") } # 使用S3原生路径 parquet_path = "s3://your-bucket/output/scaled.parquet" with make_batch_reader( parquet_path, num_epochs=4, shuffle_row_groups=False, storage_options=storage_options ) as train_reader: # 后续训练逻辑不变 train_ds = make_petastorm_dataset(train_reader).unbatch().map(lambda x: (tf.convert_to_tensor(x))).batch(2) for ele in train_ds: tensor = tf.reshape(ele,(2,1,15)) model.fit(tensor,tensor)
3. 简化认证方式(适配挂载/Unity Catalog存储)
如果目标存储已通过Databricks挂载或Unity Catalog配置,可尝试使用dbfs:/路径并添加storage_options={"use_dask": False},绕过手动认证(仅部分场景生效):
with make_batch_reader( "dbfs:/output/scaled.parquet", num_epochs=4, shuffle_row_groups=False, storage_options={"use_dask": False} ) as train_reader: # 后续训练逻辑不变 train_ds = make_petastorm_dataset(train_reader).unbatch().map(lambda x: (tf.convert_to_tensor(x))).batch(2) for ele in train_ds: tensor = tf.reshape(ele,(2,1,15)) model.fit(tensor,tensor)
关键注意事项
- 确保Linked Service或存储密钥拥有目标Parquet文件的读写权限
- 优先使用云存储原生路径(
abfs:///s3://)而非dbfs:/,避免路径解析异常 make_batch_reader与make_reader共享storage_options参数,可参考官方认证逻辑灵活适配
内容的提问来源于stack exchange,提问作者jimmy page
相关产品推荐
相关产品推荐

