【问题标题】:DQN understanding input and output (layer)DQN 理解输入和输出(层)
【发布时间】:2021-03-13 04:02:34
【问题描述】:

我有一个关于 DQN 的输入和输出(层)的问题。

例如

两个点:P1(x1, y1) 和 P2(x2, y2)

P1 必须走向 P2

我有以下信息:

  • 当前位置 P1 (x/y)
  • 当前位置 P2 (x/y)
  • 到 P1-P2 的距离 (x/y)
  • P1-P2 方向 (x/y)

P1 有 4 种可能的动作:

  • 向上
  • 向下
  • 左
  • 对

如何设置输入和输出层?

  • 4 个输入节点
  • 4 个输出节点

正确吗? 我与输出有什么关系? 我得到了 4 个数组,每个数组有 4 个值作为输出。 对输出做 argmax 是否正确?

编辑:

输入/状态:

# Current position P1
state_pos = [x_POS, y_POS]
state_pos = np.asarray(state_pos, dtype=np.float32)
# Current position P2
state_wp = [wp_x, wp_y]
state_wp = np.asarray(state_wp, dtype=np.float32)
# Distance P1 - P2 
state_dist_wp = [wp_x - x_POS, wp_y - y_POS]
state_dist_wp = np.asarray(state_dist_wp, dtype=np.float32)
# Direction P1 - P2
distance = [wp_x - x_POS, wp_y - y_POS]
norm = math.sqrt(distance[0] ** 2 + distance[1] ** 2)
state_direction_wp = [distance[0] / norm, distance[1] / norm]
state_direction_wp = np.asarray(state_direction_wp, dtype=np.float32)
state = [state_pos, state_wp, state_dist_wp, state_direction_wp]
state = np.array(state)

网络:

def __init__(self):
    self.q_net = self._build_dqn_model()
    self.epsilon = 1 

def _build_dqn_model(self):
    q_net = Sequential()
    q_net.add(Dense(4, input_shape=(4,2), activation='relu', kernel_initializer='he_uniform'))
    q_net.add(Dense(128, activation='relu', kernel_initializer='he_uniform'))
    q_net.add(Dense(128, activation='relu', kernel_initializer='he_uniform'))
    q_net.add(Dense(4, activation='linear', kernel_initializer='he_uniform'))
    rms = tf.optimizers.RMSprop(lr = 1e-4)
    q_net.compile(optimizer=rms, loss='mse')
    return q_net

def random_policy(self, state):
    return np.random.randint(0, 4)

def collect_policy(self, state):
    if np.random.random() < self.epsilon:
        return self.random_policy(state)
    return self.policy(state)

def policy(self, state):
    # Here I get 4 arrays with 4 values each as output
    action_q = self.q_net(state)

【问题讨论】:

    标签: python deep-learning reinforcement-learning q-learning dqn


    【解决方案1】:

    在第一个密集层中添加input_shape=(4,2) 会导致输出形状为(None, 4, 4)。 用以下方式定义 q_net 即可解决:

    q_net = Sequential()
    q_net.add(Reshape(target_shape=(8,), input_shape=(4,2)))
    q_net.add(Dense(128,  activation='relu', kernel_initializer='he_uniform'))
    q_net.add(Dense(128, activation='relu', kernel_initializer='he_uniform'))
    q_net.add(Dense(128, activation='relu', kernel_initializer='he_uniform'))
    q_net.add(Dense(4, activation='linear', kernel_initializer='he_uniform'))
    rms = tf.optimizers.RMSprop(lr = 1e-4)
    q_net.compile(optimizer=rms, loss='mse')
    return q_net
    

    这里,q_net.add(Reshape(target_shape=(8,), input_shape=(4,2))) 将 (None, 4, 2) 输入重塑为 (None, 8) [这里,None 表示批处理形状]。

    为了验证,打印q_net.output_shape,它应该是(None, 4) [而在前面的例子中是(None, 4, 4)]。

    你还需要做一件事。回想一下input_shape 没有考虑批处理形状。我的意思是,input_shape=(4,2) 期望输入形状 (batch_shape, 4, 2)。通过打印q_net.input_shape 来验证它,它应该输出(None, 4, 2)。现在,您需要做的是 - 在您的输入中添加一个批次维度。您只需执行以下操作:

    state_with_batch_dim = np.expand_dims(state,0)
    

    并将state_with_batch_dim 作为输入传递给 q_net。例如,您可以像 policy(np.expand_dims(state,0)) 一样调用您编写的 policy 方法,并获得维度为 (batch_shape, 4) [在本例中为 (1,4)] 的输出。

    以下是您最初问题的答案:

    1. 您的输出层应该有 4 个节点(单元)。
    2. 您的第一个密集层不一定必须有 4 个节点(单元)。如果您考虑Reshape 层,则节点或单元的概念不适合那里。您可以将Reshape 层视为一个占位符,它采用形状为 (None, 4, 2) 的张量并输出形状为 (None, 8) 的重新调整的张量。
    3. 现在,您应该得到形状为 (None, 4) 的输出 - 这 4 个值代表 4 个相应动作的 q 值。无需在此处执行argmax 即可找到 q 值。

    【讨论】:

      【解决方案2】:

      向 DQN 提供一些关于它当前所面临方向的信息也很有意义。您可以将其设置为 (Current Pos X, Current Pos Y, X From Goal, Y From Goal, Direction)。

      输出层应该按照您确定的顺序(上、左、下、右)。 Argmax 层适用于该问题。确切的代码取决于您是否使用 TF / Pytorch。

      【讨论】:

      • 感谢您的回答。我正在使用 TF。我不明白我得到的输出。 4 个数组,因为 4 个输出节点和 4 个可能的动作,对吧?但是为什么我在每个数组中得到 4 个值呢?
      • 你使用的神经网络是什么形状的?
      • 1 个输入层,4 个节点,2 个密集层,每个 128 个节点,1 个输出层,4 个节点
      • 我很难理解为什么你会得到那个输出层,很抱歉。我主要使用 Pytorch。
      • 没问题。通常有 4 个输出,也就是 4 个动作,我会得到 4 个 q 值,对吧?
      猜你喜欢
      • 1970-01-01
      • 1970-01-01
      • 2019-04-27
      • 1970-01-01
      • 2018-04-11
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 2017-09-19
      相关资源
      最近更新 更多