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

TensorFlow AI控制方块移动异常:仅能向右/向下移动求助

问题:TensorFlow自主收集圆点AI仅能向右/向下移动的解决方法

问题描述

我正在学习TensorFlow,尝试开发一个可自主学习收集随机生成圆点的AI,但运行代码时发现方块仅能向右或向下移动。已调整相关参数,确认该移动是模型输出导致的,特此求助。


核心问题分析

  1. 输入输出逻辑完全错位:当前模型输入是方块的绝对坐标,但训练数据用的是圆点位置,输出是无意义的0/1序列,模型根本不知道要根据目标位置调整移动方向。
  2. 激活函数与移动逻辑不匹配:输出层用tanh(输出范围[-1,1]),但初始训练没有教模型输出负向值(对应向左/向上),导致模型只会输出非负值。
  3. 缺乏有效的强化学习反馈:当前是一次性监督训练,没有在游戏过程中根据"吃到圆点加分、撞墙扣分"的奖励信号更新模型,无法让AI学习最优动作。
  4. 学习率设置错误:碰撞边界时将score(此时已重置为0)赋值给优化器学习率,导致模型参数完全无法更新。

解决方案

1. 修正输入特征

模型输入应该是方块与最近圆点的相对坐标,而非方块的绝对位置,这样模型能明确目标方向:

# 找到最近的圆点
if len(dot_positions) > 0:
    nearest_dot = dot_positions[np.argmin(np.linalg.norm(dot_positions - np.array([square_x, square_y]), axis=1))]
    inputs = np.array([nearest_dot[0] - square_x, nearest_dot[1] - square_y])  # 相对坐标
else:
    inputs = np.array([0, 0])

2. 重构模型输出层

改为输出4个方向(上、下、左、右)的概率,用softmax激活,对应明确的移动逻辑:

model = tf.keras.models.Sequential([
    tf.keras.layers.Dense(128, input_shape=(2,), activation='relu'),
    tf.keras.layers.Dense(4, activation='softmax')  # 4个方向:上、下、左、右
])

对应移动逻辑:

direction = np.argmax(prediction)
if direction == 0:  # 上
    square_y -= speed
elif direction == 1:  # 下
    square_y += speed
elif direction == 2:  # 左
    square_x -= speed
elif direction == 3:  # 右
    square_x += speed

3. 加入强化学习反馈机制

移除无效的初始model.fit,改为在游戏过程中收集经验(状态、动作、奖励、下一个状态),根据奖励信号更新模型:

  • 吃到圆点:奖励+10
  • 撞墙:奖励-20
  • 无有效动作:奖励-1(鼓励AI主动移动)

4. 修复学习率设置

删除model.optimizer.lr.assign(score),改用固定合理的学习率:

model.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=1e-3), 
              loss='sparse_categorical_crossentropy', 
              metrics=['accuracy'])

修改后的完整代码示例

import tensorflow as tf
import numpy as np
import pygame

pygame.init()

SCREEN_SIZE = (400, 400)
screen = pygame.display.set_mode(SCREEN_SIZE)
pygame.display.set_caption("TensorFlow Dot Collector")
DOT_COLOR = (0, 0, 255)
SQUARE_COLOR = (255, 0, 0)
DOT_RADIUS = 10
NUM_DOTS = 10
SQUARE_SIZE = 30
square_x = SCREEN_SIZE[0] // 2
square_y = SCREEN_SIZE[1] // 2
game_loop = True
score = 0
speed = 5
learning_rate = 1e-3

# 初始化圆点位置
dot_positions = np.random.randint(0, SCREEN_SIZE[0], (NUM_DOTS, 2))
old_dotpositions = dot_positions.copy()

font = pygame.font.Font(None, 36)

# 重构模型:输入是相对坐标,输出4个方向概率
model = tf.keras.models.Sequential([
    tf.keras.layers.Dense(128, input_shape=(2,), activation='relu'),
    tf.keras.layers.Dense(4, activation='softmax')
])

model.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=learning_rate),
              loss='sparse_categorical_crossentropy',
              metrics=['accuracy'])

# 经验回放缓冲区
experience_buffer = []
buffer_size = 1000
batch_size = 32

while game_loop:
    for event in pygame.event.get():
        if event.type == pygame.QUIT:
            game_loop = False

    # 获取当前状态:方块与最近圆点的相对坐标
    if len(dot_positions) > 0:
        distances = np.linalg.norm(dot_positions - np.array([square_x, square_y]), axis=1)
        nearest_idx = np.argmin(distances)
        nearest_dot = dot_positions[nearest_idx]
        current_state = np.array([nearest_dot[0] - square_x, nearest_dot[1] - square_y])
    else:
        current_state = np.array([0, 0])
        # 重新生成圆点
        dot_positions = np.random.randint(0, SCREEN_SIZE[0], (NUM_DOTS, 2))
        old_dotpositions = dot_positions.copy()

    # 模型预测动作
    prediction = model.predict(current_state.reshape(1, 2), verbose=0)
    direction = np.argmax(prediction)

    # 执行动作,记录下一个状态和奖励
    reward = -1  # 每步基础惩罚,鼓励快速移动
    prev_x, prev_y = square_x, square_y
    if direction == 0:  # 上
        square_y -= speed
    elif direction == 1:  # 下
        square_y += speed
    elif direction == 2:  # 左
        square_x -= speed
    elif direction == 3:  # 右
        square_x += speed

    # 检查碰撞圆点
    collision = False
    new_dot_positions = dot_positions.copy()
    for i, dot_pos in enumerate(dot_positions):
        if np.linalg.norm(np.array([square_x, square_y]) - dot_pos) < DOT_RADIUS + SQUARE_SIZE // 2:
            score += 5
            reward = 10  # 吃到圆点奖励
            new_dot_positions = np.delete(new_dot_positions, i, axis=0)
            collision = True
    dot_positions = new_dot_positions

    # 检查撞墙
    if square_x < 0 or square_y < 0 or square_x > SCREEN_SIZE[0] or square_y > SCREEN_SIZE[1]:
        square_x, square_y = prev_x, prev_y  # 回退位置
        score -= 10
        reward = -20  # 撞墙惩罚

    # 获取下一个状态
    if len(dot_positions) > 0:
        distances = np.linalg.norm(dot_positions - np.array([square_x, square_y]), axis=1)
        nearest_idx = np.argmin(distances)
        nearest_dot = dot_positions[nearest_idx]
        next_state = np.array([nearest_dot[0] - square_x, nearest_dot[1] - square_y])
    else:
        next_state = np.array([0, 0])

    # 存入经验缓冲区
    experience_buffer.append((current_state, direction, reward, next_state))
    if len(experience_buffer) > buffer_size:
        experience_buffer.pop(0)

    # 训练模型(当缓冲区有足够数据时)
    if len(experience_buffer) >= batch_size:
        batch = np.random.choice(len(experience_buffer), batch_size, replace=False)
        states = np.array([experience_buffer[i][0] for i in batch])
        actions = np.array([experience_buffer[i][1] for i in batch])
        rewards = np.array([experience_buffer[i][2] for i in batch])
        next_states = np.array([experience_buffer[i][3] for i in batch])

        # 计算目标Q值(简单DQN)
        next_predictions = model.predict(next_states, verbose=0)
        target_q = rewards + 0.9 * np.max(next_predictions, axis=1)  # 折扣因子0.9

        # 更新模型
        model.train_on_batch(states, actions)

    # 绘制画面
    screen.fill((255, 255, 255))
    for dot_position in dot_positions:
        pygame.draw.circle(screen, DOT_COLOR, dot_position, DOT_RADIUS)
    pygame.draw.rect(screen, SQUARE_COLOR, 
                     (square_x - SQUARE_SIZE // 2, square_y - SQUARE_SIZE // 2, 
                      SQUARE_SIZE, SQUARE_SIZE))
    score_text = font.render(f"Score: {score}", True, (0, 0, 0))
    screen.blit(score_text, (10, 10))
    pygame.display.update()

pygame.quit()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 17:45:48