ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

XIRL 跨实体逆强化学习代码库实战指南:自监督视频预训练与基于表征的强化学习奖励设计

XIRL 跨实体逆强化学习代码库实战指南:自监督视频预训练与基于表征的强化学习奖励设计 人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载XIRLCross-embodiment Inverse Reinforcement Learning是 CoRL 2021 收录论文的官方开源实现本仓库将整个工作拆解为「在视频数据上做自监督预训练」与「把学到的表征当作奖励函数做下游强化学习」两个通用环节。本文以xirl/目录为核心完整梳理其环境搭建、数据集准备、配置体系、实验复现命令以及扩展机制并结合源码讲解 TCC 时序循环一致性、帧采样器、评估器与 SAC 策略训练等底层实现帮助你直接复现论文结果或在其基础上实现自己的预训练算法与奖励设计。OverviewXIRL 解决什么问题XIRL 代码库服务于论文XIRL: Cross-embodiment Inverse Reinforcement LearningKevin Zakka、Andy Zeng、Pete Florence、Jonathan Tompson、Jeannette Bohg、Debidatta DwibediCoRL 2021。从 xirl/README.md 的定位看它是一套通用库核心能力有二自监督预训练self-supervised pretraining在无标签视频数据上学到跨实体的视觉表征下游强化学习把预训练表征embedding作为稠密奖励函数训练 RL 策略实现跨实体cross-embodiment的模仿学习即用一类机械臂的演示去指导另一类机械臂完成任务。仓库同时包含论文复现所需的模型、训练脚本和配置文件。为了提升代码模块化与可扩展性作者还配套发布了两个独立库x-magical面向跨实体模仿的 Gym 类基准MAGICAL 的扩展和torchkit包含日志、模型 checkpoint 等 PyTorch 样板工具的轻量库两者均通过 xirl/requirements.txt 以 git 依赖的形式引入。Setup环境安装与依赖清单仓库开发环境为Python 3.8 Miniconda。安装步骤如下# Clone 并进入 xirl 目录。 git clone gitgithub.com:google-research/google-research.git --depth1 cd google-research/xirl # 创建并激活 conda 环境。 conda create -n xirl python3.8 conda activate xirl # 安装依赖。 pip install -r requirements.txtxirl/requirements.txt 中锁定了关键版本包括深度学习torch1.7.1、torchvision0.8.2、tensorflow-cpu2.6.0用于 TensorBoard 与部分基线、ml-collections0.1.0配置文件体系、scikit-learn0.24.1Kendalls Tau 等评估指标环境与数据gym0.17.*、pymunk5.6.0x-MAGICAL 物理仿真、albumentations0.5.2数据增强、gdown4.4.0数据集下载工具imageio/imageio-ffmpeg视频读写、tqdm、absl-py、numpy、scipy、pandas、protobuf~3.19.0两个 git 依赖torchkitv0.0.2与x-magicalv0.0.2。需要注意的是requirements.txt中锁定的 PyTorch 版本较旧在更新的 Python 版本或新硬件如较新 CUDA上安装时可能需要调整版本这是该仓库代码发布时间较早CoRL 2021 前后所带来的限制。DatasetsX-MAGICAL 演示数据集准备X-MAGICAL运行仓库自带的 bash 脚本即可下载 X-MAGICAL 基准的演示数据集bash scripts/download_xmagical_dataset.sh从 scripts/download_xmagical_dataset.sh 源码可以看到脚本通过 Google Drive 文件 ID1VdMRYu0Y-ep_vq28hW2n0UZaow2iaW1i调用gdown下载xmagical.zip解压后默认存放于/tmp/xirl/datasets/xmagical。你可以自由修改保存位置但必须同步更新预训练配置文件中的config.data.root见 base_configs/pretrain.py。X-REAL论文中的真实世界数据集 X-REAL 在 README 发布时尚未开放will be released as soon as it gets approval仓库当前并未包含该数据集请以仓库实际状态为准。Code Navigation代码库结构与两大核心脚本XIRL 的顶层逻辑依赖两个通用 Python 脚本全部实验都由 ml_collections 以「配置文件参数化脚本」的方式驱动脚本职责默认基础配置pretrain.py表征自监督预训练base_configs/pretrain.pytrain_policy.py基于表征或环境奖励的强化学习base_configs/rl.py所有实验必须使用继承自base_configs/基础配置的配置文件预训练实验继承base_configs/pretrain.pyRL 实验继承base_configs/rl.py。其余代码组织如下configs/CoRL 论文复现使用的全部配置文件均继承自base_configs/xirl/预训练核心代码库模型、损失、训练器、评估器、数据集等sac/Soft Actor-Critic 核心实现改编自 pytorch_sacscripts/杂项 bash 脚本数据集下载、多 GPU RL 启动。配置体系深度解析预训练基础配置 base_configs/pretrain.py该文件是预训练实验的默认值模板完整字段如下通用实验参数参数默认值说明root_dir/tmp/xirl/pretrain_runs/实验保存根目录seed1RNG 种子设为none可禁用cudnn_deterministicFalsecuDNN 确定性影响可复现性cudnn_benchmarkTrue是否允许 cuDNN 自动选择最优卷积算法algorithmtcc预训练算法对应factory.py中TRAINERS字典的键logging_frequency100TensorBoard 日志间隔步数checkpointing_frequency200checkpoint 保存间隔步数数据集参数config.dataroot数据集根目录绝对路径默认/tmp/xirl/datasets/xmagical/batch_sizemini-batch 大小注意它只代表一个 batch 中加载帧的视频数因为每个视频会采样多个帧序列实际有效 batch 更大pretrain_action_class/downstream_action_class选择参与预训练 / 下游评估的动作类别留空则加载全部类别max_vids_per_class每类最多视频数默认-1不限制用于「按演示数量研究样本复杂度」的实验pretraining_video_sampler视频 batch 的采样方式random随机跨类采样或same_class同一 batch 只取同一类文件夹的视频。帧采样参数config.frame_samplerimage_ext视频文件夹中图片的通配符通常为*.jpg或*.pngstrategy采样策略all/strided/variable_strided/uniform/uniform_with_positives/last_and_randoms/window见下文 FRAME_SAMPLERSnum_frames_per_sequence每段视频采样帧数默认15num_context_frames/context_stride每个采样帧附带的上下文帧数及其步长供 3D 卷积类模型使用各策略子配置all_sampler.stride1、strided_sampler.stride3 / offsetTrue、uniform_sampler.offset0、window_sampler无额外参数。数据增强参数config.data_augmentationimage_size训练分辨率默认(112, 112)train_transforms训练集增强列表默认[random_resized_crop, color_jitter, grayscale, gaussian_blur]normalize被注释掉顺序有讲究若启用normalize应放在最后eval_transforms评估集增强默认[global_resize]。在 xirl/xirl/factory.py 的TRANSFORMS字典中可以查到每个增强名称对应的 albumentations 参数例如random_resized_crop默认scale(0.8, 1.0)、color_jitter默认brightness0.4, contrast0.4, hue0.1, saturation0.1, p0.8、gaussian_blur默认blur_limit(13, 13), sigma_limit(1.0, 2.0), p0.2、grayscale概率p0.2。create_transform还支持以name::{key: value}形式内联覆盖参数。评估参数config.evalval_iters下游 dataloader 迭代次数None表示评估整个 dataloadereval_frequency两次评估之间的步数默认500downstream_task_evaluators按顺序依次运行的下游任务评估器默认[reward_visualizer, kendalls_tau]distanceembedding 空间的距离度量cosine或sqeuclidean需与损失计算一致各评估器子配置kendalls_tau.stride3、reward_visualizer.num_plots2、cycle_consistency.stride1、nearest_neighbour_visualizer.num_videos4、embedding_visualizer.num_seqs2、reconstruction_visualizer.num_frames2。模型参数config.modelmodel_type默认resnet18_linear可选项见MODELS字典resnet18_linear、resnet18_classifier、resnet18_features、resnet18_linear_aeembedding_sizeembedding 维度默认32normalize_embeddings是否归一化 embedding默认Falselearnable_temp是否可学习温度参数默认False。损失参数config.lossTCC 损失stochastic_matchingFalse、loss_typeregression_mse、cycle_length2、label_smoothing0.1、softmax_temperature0.1、normalize_indicesTrue、variance_lambda0.001、huber_delta0.1、similarity_typel2TCN 损失pos_radius1、neg_radius4、num_pairs2、margin1.0、temperature0.1LIFS 损失temperature1.0。优化器参数config.optimtrain_max_iters最大训练迭代数默认4_000weight_decayL2 正则默认1e-4lr学习率默认1e-5。优化器固定为 Adam见 factory.py 的optim_from_config。RL 基础配置 base_configs/rl.py该文件定义 SAC 策略训练的默认值。其中obs_dim、action_dim、action_range为占位符会在运行时由 train_policy.py 读取 gym 环境后动态填充env.observation_space.shape[0]、env.action_space.shape[0]、action space 的 min/max随后将配置冻结为FrozenConfigDict落盘。关键参数包装器action_repeat1、frame_stack3reward_wrapper.pretrained_path预训练实验路径为空则用环境奖励、reward_wrapper.typedistance_to_goal或goal_classifier对应 rl_xmagical_learned_reward.py 中根据预训练算法自动选择奖励类型训练参数num_train_steps75_000、replay_buffer_capacity1_000_000、num_seed_steps5_000、num_eval_episodes50、eval_frequency5_000、checkpoint_frequency50_000、log_frequency10_000、save_videoTrueSAC 参数discount0.99、init_temperature0.1、各网络lr1e-4、critic_tau0.005、critic_target_update_frequency2、actor_update_frequency1、batch_size1024、learnable_temperatureTrueActor/Critic 均为hidden_dim1024、hidden_depth2Actor 的log_std_bounds[-5, 2]。不同实体的训练步数由 configs/constants.py 的XMAGICALTrainingIterations给出longstick75_000、mediumstick250_000、shortstick500_000、gripper500_000同文件还定义了四种实体EMBODIMENTS {shortstick, mediumstick, longstick, gripper}、五种算法ALGORITHMS {xirl, tcn, lifs, goal_classifier, raw_imagenet}以及实体到 Gym 环境名的映射如SweepToTop-Longstick-State-Allo-TestLayout-v0。实验复现运行论文全部核心实验仓库根目录下的实验启动脚本均封装了对pretrain.py/train_policy.py的调用先--help查看参数再运行。核心脚本同实体设置论文 5.1 节python pretrain_xmagical_same_embodiment.py --help python rl_xmagical_learned_reward.py --help跨实体设置论文 5.2 节python pretrain_xmagical_cross_embodiment.py --help python rl_xmagical_learned_reward.py --help环境奖励 RLbaselinepython rl_xmagical_env_reward.py --help交互式奖励可视化论文 5.4 节python interact_reward.py --help启动脚本内部机制pretrain_xmagical_same_embodiment.py 与 pretrain_xmagical_cross_embodiment.py二者都通过--algo枚举ALGORITHMS和可选--embodiment指定训练目标内部用ALGO_TO_CONFIG把算法名映射到 configs/xmagical/pretraining/ 下的配置文件tcc.py/lifs.py/tcn.py/classifier.py/imagenet.py。区别在于同实体把该实体同时设为预训练与下游类别--config.data.pretrain_action_class/downstream_action_class均为该实体跨实体则用「除目标实体外的全部实体」训练trainable_embs tuple(EMBODIMENTS - set([embodiment]))从而检验表征的泛化能力。两个脚本在预训练结束后都会调用 compute_goal_embedding.py 计算并存储平均 goal embeddinggoal_classifier基线除外并把实验元数据写入metadata.yaml供后续 RL 阶段读取。rl_xmagical_learned_reward.py读取预训练实验目录下的metadata.yaml根据algo决定奖励类型goal_classifier或distance_to_goal按实体映射出环境名并通过--config.reward_wrapper.pretrained_path把预训练模型注入 RL 奖励包装器--seeds默认[0, 5]表示运行 seed 04每个 seed 以子进程并行启动。rl_xmagical_env_reward.py以稀疏环境奖励训练作为对照 baselineRL 训练步数按实体从XMAGICALTrainingIterations自动配置见 configs/xmagical/rl/env_reward.py。杂项脚本# 可视化 dataloader用于帧采样器调试 python debug_dataset.py --help # 用预训练模型计算 goal embedding python compute_goal_embedding.py --help # 快速多 GPU RL 训练环境奖励 bash scripts/launch_rl_multi_gpu.sh训练主循环源码解读pretrain.py 的主循环展示了完整训练流程校验配置validate_config位于 base_configs/init.py→ 构建实验目录setup_experiment→ 设置设备与 RNG 种子 → 通过 xirl/common.py 的get_factories加载模型、优化器、预训练/下游 dataloader、trainer 与评估管理器 → 创建 CheckpointManager支持--resume→ 逐迭代训练并周期性执行 TensorBoard 日志、预训练验证损失评估、下游任务评估与 checkpoint 保存。train_policy.py 则展示了 RL 循环动态填充 SAC 的 obs/action 维度 → 构建训练/评估两个环境评估环境可保存 rollout 视频→ 前num_seed_steps步用随机动作填充回放缓冲区 → 之后policy.update更新 SAC → 按eval_frequency评估num_eval_episodes个 episode 并记录平均指标。值得注意的细节是当配置了reward_wrapper.pretrained_path时回放缓冲区会额外插入env.render(modergb_array)的 RGB 帧供学习式奖励包装器使用。核心实现原理从工厂到训练器工厂注册表 xirl/xirl/factory.py整个代码库的扩展点集中体现在几个注册字典上配置中的字符串键与类一一对应TRANSFORMS图像增强名称 → albumentations 调用默认参数见上文FRAME_SAMPLERS帧采样策略 → 采样器类AllSampler、StridedSampler、VariableStridedSampler、UniformSampler、UniformWithPositivesSampler、LastFrameAndRandomFrames、WindowSamplerVIDEO_SAMPLERS视频采样 →RandomBatchSampler/SameClassBatchSampler/SameClassBatchSamplerDownstreamMODELS模型类型 → 网络类ResNet18 系列TRAINERS算法名 → 训练器类tcc→TCCTrainer、lifs→LIFSTrainer、tcn→TCNTrainer、goal_classifier→GoalFrameClassifierTrainerEVALUATORS评估器名 → 评估器类Kendalls Tau、双向/三向循环一致性、NN 可视化、奖励可视化、embedding 可视化、重建可视化。evaluator_from_config会按配置中的downstream_task_evaluators列表逐个实例化评估器并包进EvalManagerdataset_from_config根据downstream标志决定是创建统一数据集还是按动作类别拆分的下游数据集字典。TCC 训练器 xirl/xirl/trainers/tcc.pyTCCTemporal Cycle Consistency时序循环一致性是 XIRL 论文默认的预训练算法其实现继承自 xirl/trainers/base.py 的Trainer。TCCTrainer在__init__中把配置里loss.tcc.*的全部超参读入成员变量compute_loss中从 batch 提取帧索引frame_idxs与视频长度video_len动态计算 cycle 数量batch_size * num_cc_frames最终调用 xirl/losses.py 的compute_tcc_loss把stochastic_matching、normalize_embeddings、loss_type、similarity_type、cycle_length、temperature、label_smoothing、variance_lambda、huber_delta等参数全部传入。这也是 README 推荐阅读的「如何实现自己的自监督预训练算法」的参考范本。Extending XIRL三种典型扩展方式README 以 QA 形式给出了扩展指引结合 factory.py 的注册机制可归纳为如何实现自己的自监督预训练算法继承xirl.trainers.base.Trainer实现__init__与compute_loss两个方法参考 xirl/trainers/tcc.py然后把新算法类注册进 factory.py 的TRAINERS字典即可通过配置项algorithm一键切换。如何修改 dataloader 中的帧采样方式在 xirl/frame_samplers.py 中创建自己的采样器类并注册进FRAME_SAMPLERS字典再通过config.frame_sampler.strategy指定。如何增加额外的预训练评估指标继承 xirl/evaluators/base.py 的Evaluator类注册进EVALUATORS字典并加入config.eval.downstream_task_evaluators列表。现有的定性与定量指标循环一致性、Kendalls Tau、各类可视化均位于 xirl/evaluators/ 目录可直接作为参考。引用如果你在研究中使用了本代码库建议引用论文inproceedings{zakka2021xirl, author {Zakka, Kevin and Zeng, Andy and Florence, Pete and Tompson, Jonathan and Bohg, Jeannette and Dwibedi, Debidatta}, title {XIRL: Cross-embodiment Inverse Reinforcement Learning}, booktitle {Proceedings of the 5th Conference on Robot Learning (CoRL)}, year {2021}, }总结XIRL 仓库的价值在于把「跨实体逆强化学习」拆解为清晰的两段式流水线第一阶段用 TCC/TCN/LIFS 等自监督算法在演示视频上学取通用表征第二阶段用该表征构造稠密奖励训练 SAC 策略。通过 base_configs/ 继承式配置、factory.py 注册表以及 Trainer/Evaluator 抽象基类你可以低成本地替换算法、采样器、模型与评估指标将跨实体奖励学习复用到自己的机器人或仿真任务中。需要注意的是仓库锁定的是 2021 年前后的依赖版本且 X-REAL 真实数据集当时尚未开放动手前应结合自身环境对依赖与数据作相应调整。赞分享人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载相关推荐基于 TRL 的 GRPO 强化学习后训练实战指南从奖励函数设计到生产部署基于 TRL 的 GRPO 强化学习后训练实战指南从奖励函数设计到生产部署 GRPOGroup Relative Policy Optimization组AI 技能人工智能大模型深度学习NoRML 无奖励元学习Google Research 强化学习 MAML 实现与训练评估实战指南NoRML 无奖励元学习Google Research 强化学习 MAML 实现与训练评估实战指南 导读 NoRMLNo Reward Meta Learn人工智能深度学习NLP计算机视觉强化学习GRPO 与 RLVR 训练实战基于 TRL 的可验证奖励强化学习配方、奖励门禁与变体选型agents24 llm-finetuning 插件GRPO 与 RLVR 训练实战基于 TRL 的可验证奖励强化学习配方、奖励门禁与变体选型agents24 llm finetuning 插件 导读 本文AI 插件AI 技能开发工具创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表