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

如何在TorchRL中定义依赖环境状态的Q-learning策略?

问题:TorchRL中适配动态合法动作子集的Q-learning策略定义

我正在训练一款扑克变体卡牌游戏的模型。每个玩家回合的合法动作是其手牌的子集(可空)。我将当前手牌和动作编码为52维布尔向量:

hand = torch.zeros(self.batch_size + (52,), dtype=torch.bool)

动作规格定义如下:

self.action_spec = BinaryDiscreteTensorSpec(52, dtype=torch.bool)

玩家的合法动作列表可从自定义环境中获取,比如下面是简化错误检查后的rand_action方法:

def rand_action(self, tensordict: TensorDictBase | None = None) -> TensorDictBase:
    if tensordict is not None:
        shape = tensordict.shape
    elif not self.batch_size:
        shape = torch.Size([])
    moves = []
    for i in range(self.total_batch_size()):
        game = self.games[i]
        game_moves = game.getMoves()
        game_move = game_moves[np.random.randint(len(game_moves))]
        moves.append([bool(i in game_move) for i in range(52)])
    if len(self.batch_size) == 0:
        r = TensorDict({"action": torch.tensor(moves[0], dtype=torch.bool)}, 
                           batch_size=self.batch_size)
    else:
        r = TensorDict({"action": torch.tensor(moves, dtype=torch.bool)}, batch_size=self.batch_size)
    if tensordict is None:
        return r
    tensordict.update(r)
    return tensordict

我需要定义一个State × Vect(52) → Values的Q函数,策略要从环境返回的合法动作列表中选择最优动作。但标准的QValueActor和QValueModule类似乎是从预定义动作列表中做argmax选最优动作,TorchRL的文档和示例也都依赖预定义动作空间,找不到针对“动作空间是固定集合的动态子集”的实现示例。请问如何在TorchRL中定义这类Q-learning策略?


解决方案

针对动态合法动作子集的场景,核心思路是先计算所有合法动作的Q值,再从中选取最优动作,以下是具体实现步骤:

1. 自定义QValueModule

继承TorchRL的QValueModule,适配52维布尔动作的输入格式,实现状态与动作的特征融合及Q值输出:

from torchrl.modules import QValueModule
import torch.nn as nn
import torch

class CustomQValueModule(QValueModule):
    def __init__(self, state_dim=52, action_dim=52):
        super().__init__()
        # 状态编码器:将手牌(52维布尔向量)编码为特征
        self.state_encoder = nn.Sequential(
            nn.Flatten(),
            nn.Linear(state_dim, 128),
            nn.ReLU(),
            nn.Linear(128, 64)
        )
        # 动作编码器:处理52维布尔动作向量
        self.action_encoder = nn.Sequential(
            nn.Flatten(),
            nn.Linear(action_dim, 64),
            nn.ReLU()
        )
        # 融合特征并输出Q值
        self.q_head = nn.Sequential(
            nn.Linear(64 + 64, 32),
            nn.ReLU(),
            nn.Linear(32, 1)
        )

    def forward(self, tensordict):
        # 提取状态和动作并转浮点型
        state = tensordict["hand"].float()
        action = tensordict["action"].float()
        
        state_feat = self.state_encoder(state)
        action_feat = self.action_encoder(action)
        combined_feat = torch.cat([state_feat, action_feat], dim=-1)
        q_value = self.q_head(combined_feat)
        return tensordict.set("q_value", q_value)

2. 自定义QValueActor

继承QValueActor,重写逻辑实现“仅在合法动作中选最优”:

from torchrl.modules import QValueActor

class CustomQValueActor(QValueActor):
    def __init__(self, qvalue_module, action_spec):
        super().__init__(qvalue_module, action_spec)

    def forward(self, tensordict):
        # 从tensordict中取出合法动作列表,shape=(batch_size, num_valid_moves, 52)
        valid_actions = tensordict["valid_actions"].float()
        batch_size, num_valid = valid_actions.shape[:2]

        # 复制状态到每个合法动作对应的维度,shape=(batch_size, num_valid, 52)
        state = tensordict["hand"].unsqueeze(1).repeat(1, num_valid, 1).float()

        # 构造临时张量字典,计算所有合法动作的Q值
        temp_td = tensordict.clone()
        temp_td["hand"] = state
        temp_td["action"] = valid_actions
        q_values = self.qvalue_module(temp_td)["q_value"].squeeze(-1)  # shape=(batch_size, num_valid)

        # 选出每个样本中Q值最大的动作索引
        max_idx = torch.argmax(q_values, dim=-1)

        # 根据索引提取最优动作
        best_actions = torch.gather(
            valid_actions,
            1,
            max_idx.unsqueeze(-1).unsqueeze(-1).repeat(1, 1, 52)
        ).squeeze(1).bool()

        tensordict["action"] = best_actions
        return tensordict

3. 环境适配:传递合法动作到TensorDict

修改自定义环境,在step或reset时将合法动作列表以张量形式存入tensordict:

# 以step方法为例,添加合法动作存储逻辑
def step(self, tensordict):
    # ... 原有环境逻辑 ...
    # 初始化合法动作张量容器(需匹配batch维度)
    valid_actions_list = []
    for i in range(self.total_batch_size()):
        game = self.games[i]
        valid_moves = game.getMoves()
        # 将合法动作转换为(num_valid, 52)的布尔张量
        move_tensor = torch.tensor(
            [[bool(card in move) for card in range(52)] for move in valid_moves],
            dtype=torch.bool
        )
        valid_actions_list.append(move_tensor)
    # 拼接为(batch_size, num_valid, 52)的张量(注意处理不同样本的num_valid差异,可补0对齐)
    valid_actions = torch.nn.utils.rnn.pad_sequence(valid_actions_list, batch_first=True)
    tensordict["valid_actions"] = valid_actions
    # ... 其他环境逻辑 ...
    return tensordict

4. 训练流程集成

将自定义模块与TorchRL的DQN训练框架结合:

# 初始化模块与Actor
q_module = CustomQValueModule()
actor = CustomQValueActor(q_module, env.action_spec)

# 定义DQN损失函数
from torchrl.objectives import DQNLoss
loss_fn = DQNLoss(qvalue_module=q_module, action_space=env.action_spec)

# 训练循环(简化版)
optimizer = torch.optim.Adam(q_module.parameters(), lr=1e-3)
num_epochs = 1000
for _ in range(num_epochs):
    # 收集训练数据
    data = env.collect(actor, max_steps=1000)
    # 计算损失并更新参数
    loss = loss_fn(data)
    optimizer.zero_grad()
    loss.backward()
    optimizer.step()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 18:53:28