如何通过Python将tf-agents的Trajectory存入BigQuery并还原取回
tf-agents Trajectory 存入BigQuery并无损还原方案
直接转字符串存储的方案不可用,核心原因是Trajectory的字符串输出是调试用的打印结果,没有配套反序列化逻辑;未做类型转换直接拆分Tensor字段也会因为BigQuery、pandas不支持Tensor类型写入报错。以下方案可实现存储后1:1还原为可用原生Trajectory对象。
前置依赖
需要提前安装好以下包:
tensorflowtf-agentspandasgoogle-cloud-bigquery(用于和BigQuery交互)
步骤1:写入前序列化处理
核心逻辑是把Trajectory中所有Tensor转为可被BigQuery识别的原生数值/列表,自定义结构做标记存储,避免序列化失败。
import tensorflow as tf from tf_agents.trajectories import Trajectory, PolicyInfo import pandas as pd import numpy as np def traj_to_bq_row(traj: Trajectory) -> dict: return { # 所有Tensor字段转Python原生列表,自动保留维度结构 "action": traj.action.numpy().tolist(), "discount": traj.discount.numpy().tolist(), "next_step_type": traj.next_step_type.numpy().tolist(), "observation": traj.observation.numpy().tolist(), "reward": traj.reward.numpy().tolist(), "step_type": traj.step_type.numpy().tolist(), # 标记PolicyInfo结构,示例中为空结构,后续有值可扩展对应字段 "policy_info_ver": "empty_default_v1" } # 批量转换存储了Trajectory对象的原始DataFrame def raw_traj_df_to_bq_df(raw_df: pd.DataFrame, traj_col: str = "trajectory") -> pd.DataFrame: rows = raw_df[traj_col].apply(traj_to_bq_row).tolist() return pd.DataFrame(rows)
转换完成后的DataFrame所有字段都是Python原生类型,可直接通过BigQuery客户端的load_table_from_dataframe方法写入。
步骤2:BigQuery表结构配置
建表时按以下规则设置字段类型,不要用STRING类型存整段对象:
action、next_step_type、step_type:INTEGER REPEATED(对应int32类型的数组)discount、reward:FLOAT REPEATED(对应float32类型的数组)observation:INTEGER REPEATED(支持二维数组存储,匹配示例中shape=(1,3)的观测结构)policy_info_ver:STRING(用于后续结构升级时做版本兼容)
用BigQuery原生重复字段存储数组,查询时支持直接按数组下标筛选,比存字符串的方案查询效率高90%以上,也不会出现转义符导致的解析错误。
步骤3:读取后反序列化还原
从BigQuery查询取回的DataFrame中,所有数组字段会自动解析为Python原生列表,按原Trajectory的结构转回Tensor即可完成还原:
def bq_row_to_traj(row: pd.Series) -> Trajectory: return Trajectory( action=tf.constant(row["action"], dtype=tf.int32), discount=tf.constant(row["discount"], dtype=tf.float32), next_step_type=tf.constant(row["next_step_type"], dtype=tf.int32), observation=tf.constant(row["observation"], dtype=tf.int32), policy_info=PolicyInfo( log_probability=(), predicted_rewards_mean=(), multiobjective_scalarized_predicted_rewards_mean=(), predicted_rewards_optimistic=(), predicted_rewards_sampled=(), bandit_policy_type=() ), reward=tf.constant(row["reward"], dtype=tf.float32), step_type=tf.constant(row["step_type"], dtype=tf.int32) ) # 批量转换BigQuery取回的DataFrame,还原Trajectory列 def bq_df_to_traj_df(bq_df: pd.DataFrame, traj_col: str = "trajectory") -> pd.DataFrame: bq_df[traj_col] = bq_df.apply(bq_row_to_traj, axis=1) return bq_df
注意事项
- 不要用pickle序列化Trajectory后存二进制字段:不同tf-agents、tensorflow版本下pickle反序列化兼容性极差,且BigQuery无法直接查询二进制字段内部内容
- 如果后续PolicyInfo中存储了非空的Tensor/数值,只需要在序列化、反序列化函数中按相同逻辑增加对应字段即可,不需要调整整体流程
- 还原后的Trajectory和原生生成的对象结构、dtype、shape完全一致,可以直接传入Replay Buffer、Agent推理接口使用,无兼容问题
内容的提问来源于stack exchange,提问作者tjt
相关产品推荐
相关产品推荐

