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

如何在TorchRL滚动过程中获取自定义Gymnasium环境的额外信息?

解决TorchRL rollout中添加额外环境信息到TensorDict的问题

下面提供两种可行方案,帮你把追逃双方的绝对位置纳入rollout返回的TensorDict中:


方法1:手动实现rollout流程,主动捕获info并整合

放弃依赖内置的rollout方法,手动循环与环境交互,直接收集每一步的info数据,最后组装成包含绝对位置的TensorDict。

示例代码:

from tensordict import TensorDict
import torch

def custom_rollout(env, policy, max_steps=1000):
    obs, info = env.reset()
    # 初始化各类数据存储列表
    observations = []
    actions = []
    rewards = []
    terminateds = []
    evader_positions = []
    pursuer_positions = []
    
    with set_exploration_type(ExplorationType.MEAN), torch.no_grad():
        for _ in range(max_steps):
            action = policy(obs)["action"]
            # 执行step时直接获取info
            next_obs, reward, terminated, _, info = env.step(action)
            
            # 收集所有需要的数据
            observations.append(obs)
            actions.append(action)
            rewards.append(reward)
            terminateds.append(terminated)
            evader_positions.append(info["evader_pos"])
            pursuer_positions.append(info["pursuer-pos"])
            
            obs = next_obs
            if terminated:
                break
    
    # 转换为张量并组装成TensorDict
    eval_rollout = TensorDict({
        "observation": torch.stack(observations),
        "action": torch.stack(actions),
        "reward": torch.stack(rewards),
        "terminated": torch.stack(terminateds),
        "evader_pos": torch.stack(evader_positions),
        "pursuer_pos": torch.stack(pursuer_positions)
    }, batch_size=[])
    return eval_rollout

# 替换原有rollout调用
if i % 10 == 0:
    eval_rollout = custom_rollout(env, policy_module, max_steps=1000)

方法2:用TorchRL环境包装器自动注入info字段

如果你的环境继承自TorchRL的EnvBase,可以通过自定义包装器,在reset和step阶段自动把info中的绝对位置添加到返回的TensorDict里。

示例代码:

from torchrl.envs import EnvBase, GymWrapper

class CustomPursuitEnv(EnvBase):
    def __init__(self, env):
        super().__init__()
        self.env = GymWrapper(env)
        # 注册额外输出字段的规格
        self.add_spec("evader_pos", self.env.observation_spec["observation"].clone())
        self.add_spec("pursuer_pos", self.env.observation_spec["observation"].clone())
    
    def _reset(self, tensordict):
        reset_td = self.env._reset(tensordict)
        # 注入初始位置信息
        info = self.env.unwrapped._get_info()
        reset_td["evader_pos"] = torch.tensor(info["evader_pos"], dtype=torch.float32)
        reset_td["pursuer_pos"] = torch.tensor(info["pursuer-pos"], dtype=torch.float32)
        return reset_td
    
    def _step(self, tensordict):
        step_td = self.env._step(tensordict)
        # 注入每一步的位置信息
        info = self.env.unwrapped._get_info()
        step_td["evader_pos"] = torch.tensor(info["evader_pos"], dtype=torch.float32)
        step_td["pursuer_pos"] = torch.tensor(info["pursuer-pos"], dtype=torch.float32)
        return step_td

# 包装你的自定义Gym环境
wrapped_env = CustomPursuitEnv(your_gym_env)

# 此后使用内置rollout方法即可获取带位置的TensorDict
if i % 10 == 0:
    with set_exploration_type(ExplorationType.MEAN), torch.no_grad():
        eval_rollout = wrapped_env.rollout(1000, policy_module)

注意事项

  • 确保get_position()返回的数据可以直接转换为PyTorch张量(如列表、numpy数组)
  • 如果使用向量环境,需要调整收集逻辑适配批量数据(比如用批量处理的方式获取所有环境实例的info)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.02 03:33:30