如何在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
相关产品推荐
相关产品推荐

