[实践记录]强化学习训练实录——2048实战

在学习了一段时间基础之后,也就想进行一些工程实践

之前一直在倒立摆这个最简单的环境中进行学习和实践

有一说一是有点太简单了

上学期还在flappy bird环境中实现了一下简单的dqn和策略梯度

这不,马上要做个项目,于是就想在稍微复杂一点的环境里实现一下强化学习,顺便试试看ppo啥的效果

于是,想起了寒假整了老半天的小项目——经典游戏2048

那就按照我的实践步骤一点一点介绍吧

0.环境准备

我想弄2048的原因之一就是,Flappy_Bird和CatePole环境都在gym这个集成环境中封装好了

我想自己定义一下环境返回的奖励值(虽说实在是没啥必要吧,但是还是想自己实践一下)

于是乎,在github上找了一个简单的2048实践(对不起当时忘了star,找不到了)

它定义了简单的环境和奖励函数

就在这个环境里上dqn了

但是当时的效果并不好,就搁置了

后来就自己(用ai),用pygame库实现了一个小环境,添加了一些信息和接口。

返回信息主要是

  • 最大块数字
  • 空格数量
  • 总得分
  • 是否结束

(好吧其实跟gym返回的没啥区别)

然后就开整吧

1.强化学习改造

这部分分了三步,我最开始的想法是,先在最简单的dqn中实现一个看起来收敛的东西

然后再改策略梯度,再改ppo

那就改环境吧,把接口什么的改一改,超参数改一改

最主要的还是奖励函数和网络结构

不改不知道,一改吓一跳

网络结构这一块,首先把输入改为16个头

我把棋盘展平为一个1维向量,然后塞进网络中

接下来是两个全连接层

这一块问题不是很大

奖励函数这一块真是头大

最一开始,我觉得很简单,就是:每一步最大块取个2log,空格数乘个小系数,总得分乘一个小系数,加在一起

之后,每一局游戏把每一步的得分加在一起,得到一局的总奖励

然后就练吧,这能不收敛么

接下来吧,dqn的大方差,还有奖励函数不照气

导致loss和奖励开始极大震荡,没几轮就震荡到了几万,根本受不了

吸取了一些经验之后,发现是单步奖励太大,于是对单步奖励进行了一个clip,约束到2

还是不收敛

接下来是一个大改(在ai的帮助下)

最终的单步奖励改成了:

合并奖励:总得分,取一个log

最大奖励:上一次出现的最大块与本轮最大快做减法并取log

空格奖励:上一次的空格数减去这一次的空格数,乘一个小系数

里程碑奖励:当最大数字到达一定数字时,给一个稍微大一些的奖励

无效惩罚:输出的动作可能是无效动作,使棋盘没有运动。这时会给一个小惩罚

结束惩罚:游戏结束,会给一个惩罚

成功奖励:当达到成功目标(最大2048)时,给一个奖励

累加起来,是每一步的奖励,但还是要进行一个裁断。

使用tanh,让奖励梯度变化一下

然后,配合着ppo,训练了1.2w轮就能稳定见到512了。

2.代码

首先是2048的环境文件model.py:

from __future__ import annotations

try:
    import gymnasium as gym
except ModuleNotFoundError:  # pragma: no cover - compatibility fallback
    import gym
import gym_2048  # noqa: F401  # trigger env registration
import numpy as np
import pygame


BG_COLOR = (187, 173, 160)
EMPTY_COLOR = (205, 193, 180)
TILE_COLORS = {
    2: (238, 228, 218),
    4: (237, 224, 200),
    8: (242, 177, 121),
    16: (245, 149, 99),
    32: (246, 124, 95),
    64: (247, 96, 63),
    128: (237, 207, 114),
    256: (237, 204, 97),
    512: (237, 200, 80),
    1024: (237, 197, 63),
    2048: (237, 194, 46),
}
LIGHT_TEXT_COLOR = (249, 246, 242)
DARK_TEXT_COLOR = (119, 110, 101)

WINDOW_WIDTH = 420
WINDOW_HEIGHT = 520
CELL_SIZE = 90
CELL_GAP = 10
BOARD_LEFT = 10
BOARD_TOP = 10

# user_input -> env action
# input: 0=up, 1=down, 2=left, 3=right
# env action: 0=left, 1=up, 2=right, 3=down
USER_TO_ENV_ACTION = {0: 1, 1: 3, 2: 0, 3: 2}
ACTION_NAMES = {0: "UP", 1: "DOWN", 2: "LEFT", 3: "RIGHT"}


def to_env_action(user_action: int) -> int:
    """Map user numeric input (0..3) to gym-2048 action id."""
    if user_action not in USER_TO_ENV_ACTION:
        raise ValueError("user_action must be one of 0, 1, 2, 3")
    return USER_TO_ENV_ACTION[user_action]


def draw_board(screen, board, score_hint, step_id):
    screen.fill(BG_COLOR)

    for i in range(4):
        for j in range(4):
            value = int(board[i, j])
            left = BOARD_LEFT + j * (CELL_SIZE + CELL_GAP)
            top = BOARD_TOP + i * (CELL_SIZE + CELL_GAP)
            color = EMPTY_COLOR if value == 0 else TILE_COLORS.get(value, (60, 58, 50))
            pygame.draw.rect(
                screen, color, (left, top, CELL_SIZE, CELL_SIZE), border_radius=8
            )

            if value != 0:
                digits = len(str(value))
                font_size = max(24, 56 - digits * 8)
                font = pygame.font.SysFont("Arial", font_size, bold=True)
                text_color = DARK_TEXT_COLOR if value <= 4 else LIGHT_TEXT_COLOR
                text = font.render(str(value), True, text_color)
                text_x = left + (CELL_SIZE - text.get_width()) / 2
                text_y = top + (CELL_SIZE - text.get_height()) / 2
                screen.blit(text, (text_x, text_y))

    info_font = pygame.font.SysFont("Arial", 22, bold=True)
    info_text = info_font.render(
        f"BoardSum: {score_hint}   Step: {step_id}", True, (250, 248, 239)
    )
    screen.blit(info_text, (10, 430))
    pygame.display.flip()


class Game2048:
    """Simple interface: reset() + step(0..3) + close()."""

    def __init__(
        self, seed=42, fps=30, window_title="2048 Gymnasium", render_enabled=True
    ):
        pygame.init()
        self.window_title = str(window_title)
        self.render_enabled = bool(render_enabled)
        self.screen = None
        if self.render_enabled:
            self._ensure_window()
        self.clock = pygame.time.Clock()

        self.env = gym.make("2048-extended-v2")
        self.fps = int(fps)
        self.closed = False

        self.board = None
        self.info = {}
        self.done = False
        self.step_id = 0
        self.reset(seed=seed)

    def _ensure_window(self):
        if not pygame.display.get_init():
            pygame.display.init()
        if self.screen is None:
            self.screen = pygame.display.set_mode((WINDOW_WIDTH, WINDOW_HEIGHT))
            pygame.display.set_caption(self.window_title)

    def _process_events(self):
        if not self.render_enabled or not pygame.display.get_init():
            return
        for event in pygame.event.get():
            if event.type == pygame.QUIT:
                self.done = True
            elif event.type == pygame.KEYDOWN and event.key == pygame.K_ESCAPE:
                self.done = True

    def _compose_info(self, base_info):
        info = dict(base_info)
        info["empty_cells"] = int(np.count_nonzero(self.board == 0))
        info["max_tile"] = int(np.max(self.board))
        info["step_id"] = int(self.step_id)
        return info

    def reset(self, seed=None):
        if self.closed:
            raise RuntimeError("Game2048 has been closed.")
        try:
            reset_out = self.env.reset(seed=seed)
        except TypeError:
            reset_out = self.env.reset()

        if isinstance(reset_out, tuple) and len(reset_out) == 2:
            self.board, base_info = reset_out
        else:
            self.board, base_info = reset_out, {}

        self.done = False
        self.step_id = 0
        self.info = self._compose_info(base_info)
        self.render()
        return self.board.copy(), dict(self.info)

    def step(self, user_action):
        """
        Execute one step by user action code.

        Args:
            user_action (int): user action id in {0, 1, 2, 3}
                0=UP, 1=DOWN, 2=LEFT, 3=RIGHT.

        Returns:
            tuple[np.ndarray, float, bool, dict]:
                board (np.ndarray):
                    current board after this step, shape (4, 4), returned as a copy.
                reward (float):
                    immediate reward returned by env for this step.
                    If user closes the window/presses ESC, this returns 0.0.
                done (bool):
                    whether current episode is finished.
                    True when terminated/truncated or closed by user.
                info (dict):
                    merged runtime info for logging/training, including:
                    - end_value / max_block / is_success (from env)
                    - empty_cells / max_tile / step_id (computed here)
                    - user_action / action_name (added here)
                    - closed_by_user=True (only in early-close branch)
        """
        if self.closed:
            raise RuntimeError("Game2048 has been closed.")

        self._process_events()
        if self.done:
            info = dict(self.info)
            info["closed_by_user"] = True
            return self.board.copy(), 0.0, True, info

        env_action = to_env_action(int(user_action))
        step_out = self.env.step(env_action)
        if isinstance(step_out, tuple) and len(step_out) == 5:
            self.board, reward, terminated, truncated, base_info = step_out
        elif isinstance(step_out, tuple) and len(step_out) == 4:
            self.board, reward, done, base_info = step_out
            terminated, truncated = bool(done), False
        else:
            raise RuntimeError(
                "Unexpected env.step() return format; expected 4 or 5 values."
            )

        self.step_id += 1
        self.done = bool(terminated or truncated)
        self.info = self._compose_info(base_info)
        self.info["user_action"] = int(user_action)
        self.info["action_name"] = ACTION_NAMES[int(user_action)]

        self.render()
        if self.render_enabled:
            self.clock.tick(self.fps)
        # Return a stable 4-tuple interface for training loops.
        return self.board.copy(), float(reward), self.done, dict(self.info)

    def render(self):
        if not self.render_enabled:
            return
        self._ensure_window()
        score_hint = int(self.info.get("end_value", np.sum(self.board)))
        draw_board(self.screen, self.board, score_hint, self.step_id)

    def set_render_enabled(self, enabled: bool):
        enabled = bool(enabled)
        if enabled == self.render_enabled:
            return

        self.render_enabled = enabled
        if self.render_enabled:
            self._ensure_window()
            self.render()
        else:
            if pygame.display.get_init():
                pygame.display.quit()
            self.screen = None

    def close(self):
        if self.closed:
            return
        self.closed = True
        self.env.close()
        if pygame.display.get_init():
            pygame.display.quit()
        pygame.quit()


def run(game: Game2048, action: int):
    """Single-step helper for external loops."""
    if action not in {0, 1, 2, 3}:
        raise ValueError("action must be one of 0, 1, 2, 3")

    return game.step(action)


def play_manual():
    game = Game2048(seed=42, fps=30, window_title="2048 Manual Control")
    try:
        done = False
        while not done:
            raw = input(
                "Input action [0=up, 1=down, 2=left, 3=right, q=quit]: "
            ).strip()
            if raw.lower() in {"q", "quit", "exit"}:
                break
            if raw not in {"0", "1", "2", "3"}:
                print("Invalid input, please enter 0/1/2/3 (or q to quit).")
                continue

            done, _ = run(game, int(raw))

        print(f"Total Moves: {game.step_id}")
    finally:
        game.close()


if __name__ == "__main__":
    play_manual()
View Code

然后是训练文件:ppo_clip.py

from __future__ import annotations

import math
import time
from pathlib import Path

import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.distributions import Categorical
from torch.utils.data.sampler import BatchSampler, SubsetRandomSampler

# Hyper Parameters
BATCH_SIZE = 512
LR_POLICY = 3e-4
LR_VALUE = 1e-3
GAMMA = 0.99
LAMBDA = 0.95
EPS_CLIP = 0.2
K_EPOCHS = 4
UPDATE_TIMESTEP = 2000
NUM_EPISODES = 20000
RENDER_EVERY_EPISODES = 100
SAVE_EVERY_EPISODES = 500
CHECKPOINT_DIR = Path(__file__).resolve().parent / "checkpoints"
N_ACTIONS = 4
N_STATES = 16
REQUIRE_CUDA = False  # Set True to fail fast when CUDA is unavailable.
AUTO_RESUME_LATEST = True  # Auto-load latest checkpoint in CHECKPOINT_DIR.
MAX_TILE_MILESTONE_BONUSES = [
    (128, 0.2),
    (256, 0.4),
    (512, 0.8),
    (1024, 1.2),
    (2048, 2.0),
]
REWARD_TANH_SCALE = 2.0  # Smooth reward squashing scale.

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
if REQUIRE_CUDA and device.type != "cuda":
    raise RuntimeError(
        "CUDA is required but unavailable. Please use a CUDA-enabled PyTorch build in your conda env."
    )
if device.type == "cuda":
    torch.backends.cudnn.benchmark = True


class PolicyNet(nn.Module):
    # 定义策略网络
    """Policy network for discrete actions."""

    def __init__(self):
        """Initialize policy network layers."""
        super().__init__()
        self.fc1 = nn.Linear(N_STATES, 512)
        self.fc2 = nn.Linear(512, 512)
        self.out = nn.Linear(512, N_ACTIONS)

        nn.init.orthogonal_(self.fc1.weight, gain=np.sqrt(2))
        nn.init.orthogonal_(self.fc2.weight, gain=np.sqrt(2))
        nn.init.orthogonal_(self.out.weight, gain=0.01)
        nn.init.zeros_(self.fc1.bias)
        nn.init.zeros_(self.fc2.bias)
        nn.init.zeros_(self.out.bias)

    def forward(self, x):
        """Return action probabilities."""
        x = F.relu(self.fc1(x))
        x = F.relu(self.fc2(x))
        logits = self.out(x)
        return F.softmax(logits, dim=-1)


class ValueNet(nn.Module):
    """State-value network."""

    # 定义值函数网络

    def __init__(self):
        """Initialize value network layers."""
        super().__init__()
        self.fc1 = nn.Linear(N_STATES, 512)
        self.fc2 = nn.Linear(512, 512)
        self.out = nn.Linear(512, 1)

        nn.init.orthogonal_(self.fc1.weight, gain=np.sqrt(2))
        nn.init.orthogonal_(self.fc2.weight, gain=np.sqrt(2))
        nn.init.orthogonal_(self.out.weight, gain=1.0)
        nn.init.zeros_(self.fc1.bias)
        nn.init.zeros_(self.fc2.bias)
        nn.init.zeros_(self.out.bias)

    def forward(self, x):
        """Return state value V(s)."""
        x = F.relu(self.fc1(x))
        x = F.relu(self.fc2(x))
        return self.out(x)


class PPO:
    """PPO-Clip agent for 2048."""

    def __init__(self):
        """Initialize policy/value networks and optimizers."""
        self.device = device
        self.policy_net = PolicyNet().to(self.device)
        self.value_net = ValueNet().to(self.device)
        self.policy_optimizer = torch.optim.Adam(
            self.policy_net.parameters(), lr=LR_POLICY
        )
        self.value_optimizer = torch.optim.Adam(
            self.value_net.parameters(), lr=LR_VALUE
        )
        self.training_step = 0

    def save_checkpoint(self, path: Path, episode: int):
        """Save policy/value networks and optimizer states."""
        path.parent.mkdir(parents=True, exist_ok=True)
        torch.save(
            {
                "episode": episode,
                "training_step": self.training_step,
                "policy_net_state_dict": self.policy_net.state_dict(),
                "value_net_state_dict": self.value_net.state_dict(),
                "policy_optimizer_state_dict": self.policy_optimizer.state_dict(),
                "value_optimizer_state_dict": self.value_optimizer.state_dict(),
            },
            path,
        )

    @staticmethod
    def _move_optimizer_state_to_device(
        optimizer: torch.optim.Optimizer, device: torch.device
    ):
        """Move optimizer state tensors to target device."""
        for state in optimizer.state.values():
            for key, value in state.items():
                if torch.is_tensor(value):
                    state[key] = value.to(device)

    def load_checkpoint(self, path: Path) -> int:
        """Load checkpoint and return saved episode index."""
        checkpoint = torch.load(path, map_location=self.device)
        self.policy_net.load_state_dict(checkpoint["policy_net_state_dict"])
        self.value_net.load_state_dict(checkpoint["value_net_state_dict"])
        self.policy_optimizer.load_state_dict(checkpoint["policy_optimizer_state_dict"])
        self.value_optimizer.load_state_dict(checkpoint["value_optimizer_state_dict"])
        self.training_step = int(checkpoint.get("training_step", 0))
        self._move_optimizer_state_to_device(self.policy_optimizer, self.device)
        self._move_optimizer_state_to_device(self.value_optimizer, self.device)
        return int(checkpoint.get("episode", 0))

    @staticmethod
    def _to_state_vector(s):
        # 把2048的棋盘状态转换为一个扁平化的log2特征向量
        # 取log2是为了让数值范围更适合神经网络处理
        # 同时保留了棋盘上块的相对大小信息。0块保持为0,其他块转换为它们的log2值。
        # 老写法里直接把数值输入到网络中显然不太合适,因为2048的块数值范围很大
        # 直接输入可能导致训练不稳定。通过取log2,我们可以把数值范围压缩到更合理的范围内,同时保留了块之间的相对关系。
        """Convert board to flattened log2 feature vector."""
        if isinstance(s, tuple):
            s = s[0]
        arr = np.asarray(s, dtype=np.float32).reshape(-1)
        mask = arr > 0
        arr[mask] = np.log2(arr[mask])
        return arr

    @staticmethod
    def _to_board(s):
        """Convert state to 4x4 integer board."""
        # 这个函数是为了在计算有效动作掩码时使用的
        # 它将状态转换回4x4的整数棋盘格式
        # 以便我们可以模拟每个动作并检查它们是否会改变棋盘状态。
        if isinstance(s, tuple):
            s = s[0]
        return np.asarray(s, dtype=np.int64).reshape(4, 4)

    @staticmethod
    def _move_row_left(row: np.ndarray) -> np.ndarray:
        # 这个函数模拟了2048游戏中将一行向左移动的逻辑
        """Simulate moving one row to the left."""
        non_zero = row[row != 0]
        merged = []
        i = 0
        while i < len(non_zero):
            if i + 1 < len(non_zero) and non_zero[i] == non_zero[i + 1]:
                merged.append(int(non_zero[i] * 2))
                i += 2
            else:
                merged.append(int(non_zero[i]))
                i += 1
        if len(merged) < 4:
            merged.extend([0] * (4 - len(merged)))
        return np.asarray(merged, dtype=np.int64)

    @classmethod
    def _move_left(cls, board: np.ndarray) -> np.ndarray:
        # 这个函数模拟了将整个棋盘向左移动的逻辑
        """Simulate moving board left."""
        return np.vstack([cls._move_row_left(row) for row in board])

    @classmethod
    def _move_right(cls, board: np.ndarray) -> np.ndarray:
        # 这个函数模拟了将整个棋盘向右移动的逻辑
        """Simulate moving board right."""
        flipped = np.fliplr(board)
        moved = cls._move_left(flipped)
        return np.fliplr(moved)

    @classmethod
    def _apply_action_to_board(cls, board: np.ndarray, action: int) -> np.ndarray:
        # 这个函数根据策略网络选择的动作(上、下、左、右)来模拟棋盘状态的变化
        """Apply user action id to board simulation."""
        # user action mapping: 0=up, 1=down, 2=left, 3=right
        if action == 0:
            return cls._move_left(board.T).T
        if action == 1:
            return cls._move_right(board.T).T
        if action == 2:
            return cls._move_left(board)
        if action == 3:
            return cls._move_right(board)
        raise ValueError("action must be one of 0, 1, 2, 3")

    @classmethod
    def get_valid_action_mask(cls, state) -> np.ndarray:
        """Return mask of valid actions at current state."""
        board = cls._to_board(state)
        valid = np.zeros(N_ACTIONS, dtype=bool)
        for action in range(N_ACTIONS):
            moved = cls._apply_action_to_board(board, action)
            valid[action] = not np.array_equal(board, moved)
        return valid

    @staticmethod
    def _apply_prob_mask(probs: torch.Tensor, valid_mask: torch.Tensor) -> torch.Tensor:
        """Apply valid-action mask to probabilities and renormalize."""
        squeeze_out = False
        if probs.dim() == 1:
            probs = probs.unsqueeze(0)
            valid_mask = valid_mask.unsqueeze(0)
            squeeze_out = True

        mask = valid_mask.bool()
        masked_probs = torch.where(mask, probs, torch.zeros_like(probs))
        denom = masked_probs.sum(dim=-1, keepdim=True)
        mask_count = mask.sum(dim=-1, keepdim=True).float()

        fallback_from_mask = mask.float() / mask_count.clamp(min=1.0)
        uniform = torch.full_like(probs, 1.0 / probs.size(-1))
        fallback = torch.where(mask_count > 0, fallback_from_mask, uniform)
        probs = torch.where(
            denom > 1e-8, masked_probs / denom.clamp(min=1e-8), fallback
        )

        if squeeze_out:
            probs = probs.squeeze(0)
        return probs

    def choose_action(self, state, valid_mask=None):
        """Sample action from masked policy and return action/log_prob/value."""
        state_vec = self._to_state_vector(state)
        state_t = torch.as_tensor(
            state_vec, dtype=torch.float32, device=self.device
        ).unsqueeze(0)

        probs = self.policy_net(state_t).squeeze(0)
        if valid_mask is not None:
            mask_t = torch.as_tensor(valid_mask, dtype=torch.bool, device=self.device)
            probs = self._apply_prob_mask(probs, mask_t)

        dist = Categorical(probs)
        action = dist.sample()
        log_prob = dist.log_prob(action)
        value = self.value_net(state_t).squeeze(-1)
        return int(action.item()), float(log_prob.item()), float(value.item())

    def compute_gae(self, rewards, values, dones, next_value):
        """Compute GAE advantages and returns."""
        advantages = np.zeros(len(rewards), dtype=np.float32)
        last_advantage = 0.0
        for t in reversed(range(len(rewards))):
            if t == len(rewards) - 1:
                next_val = next_value
            else:
                next_val = values[t + 1] * (1.0 - float(dones[t]))
            delta = rewards[t] + GAMMA * next_val - values[t]
            advantages[t] = (
                delta + GAMMA * LAMBDA * (1.0 - float(dones[t])) * last_advantage
            )
            last_advantage = advantages[t]
        returns = advantages + np.asarray(values, dtype=np.float32)
        return advantages, returns

    def learn(self, buffer):
        """Run PPO clipped updates on collected rollout buffer."""
        if len(buffer["states"]) == 0:
            return None, None

        states = torch.as_tensor(
            np.asarray(buffer["states"], dtype=np.float32),
            dtype=torch.float32,
            device=self.device,
        )
        actions = torch.as_tensor(
            np.asarray(buffer["actions"], dtype=np.int64),
            dtype=torch.long,
            device=self.device,
        ).unsqueeze(1)
        old_log_probs = torch.as_tensor(
            np.asarray(buffer["log_probs"], dtype=np.float32),
            dtype=torch.float32,
            device=self.device,
        ).unsqueeze(1)
        valid_masks = torch.as_tensor(
            np.asarray(buffer["valid_masks"], dtype=np.bool_),
            dtype=torch.bool,
            device=self.device,
        )
        advantages = torch.as_tensor(
            np.asarray(buffer["advantages"], dtype=np.float32),
            dtype=torch.float32,
            device=self.device,
        ).unsqueeze(1)
        returns = torch.as_tensor(
            np.asarray(buffer["returns"], dtype=np.float32),
            dtype=torch.float32,
            device=self.device,
        ).unsqueeze(1)

        advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)

        last_policy_loss = None
        last_value_loss = None

        for _ in range(K_EPOCHS):
            sampler = BatchSampler(
                SubsetRandomSampler(range(len(states))),
                batch_size=BATCH_SIZE,
                drop_last=False,
            )
            for indices in sampler:
                batch_states = states[indices]
                batch_actions = actions[indices]
                batch_old_log_probs = old_log_probs[indices]
                batch_valid_masks = valid_masks[indices]
                batch_advantages = advantages[indices]
                batch_returns = returns[indices]

                probs = self.policy_net(batch_states)
                probs = self._apply_prob_mask(probs, batch_valid_masks)
                dist = Categorical(probs)
                log_probs = dist.log_prob(batch_actions.squeeze(1)).unsqueeze(1)

                ratios = torch.exp(log_probs - batch_old_log_probs)
                surr1 = ratios * batch_advantages
                surr2 = (
                    torch.clamp(ratios, 1 - EPS_CLIP, 1 + EPS_CLIP) * batch_advantages
                )
                policy_loss = -torch.min(surr1, surr2).mean()

                self.policy_optimizer.zero_grad()
                policy_loss.backward()
                nn.utils.clip_grad_norm_(self.policy_net.parameters(), 0.5)
                self.policy_optimizer.step()

                values = self.value_net(batch_states)
                value_loss = F.mse_loss(values, batch_returns)

                self.value_optimizer.zero_grad()
                value_loss.backward()
                nn.utils.clip_grad_norm_(self.value_net.parameters(), 0.5)
                self.value_optimizer.step()

                self.training_step += 1
                last_policy_loss = float(policy_loss.item())
                last_value_loss = float(value_loss.item())

        return last_policy_loss, last_value_loss


def new_buffer():
    """Create an empty rollout buffer."""
    return {
        "states": [],
        "actions": [],
        "log_probs": [],
        "valid_masks": [],
        "advantages": [],
        "returns": [],
    }


def find_latest_checkpoint(ckpt_dir: Path):
    """Find latest checkpoint by episode number in filename."""
    if not ckpt_dir.exists():
        return None
    candidates = list(ckpt_dir.glob("ppo_clip_ep_*.pt"))
    if not candidates:
        return None

    def episode_from_name(path: Path) -> int:
        try:
            return int(path.stem.rsplit("_", 1)[-1])
        except (ValueError, IndexError):
            return -1

    return max(candidates, key=episode_from_name)


def main():
    """Train PPO-Clip on 2048."""
    from model import Game2048

    ppo = PPO()
    game = Game2048(
        seed=42, fps=30, window_title="2048 PPO-Clip Train", render_enabled=False
    )

    print(f"Using device: {device}")
    if device.type == "cuda":
        print(
            f"CUDA device: {torch.cuda.get_device_name(0)} | "
            f"torch_cuda: {torch.version.cuda}"
        )
    else:
        print(
            "CUDA unavailable in current Python env; training will run on CPU. "
            "Install CUDA-enabled PyTorch in this conda env to use GPU."
        )
    print("\nTraining with PPO-Clip...")
    print(
        f"PPO Parameters: clip_epsilon={EPS_CLIP}, k_epochs={K_EPOCHS}, batch_size={BATCH_SIZE}"
    )
    time.sleep(2)

    buffer = new_buffer()
    timestep = 0
    last_policy_loss = None
    last_value_loss = None
    start_episode = 0

    if AUTO_RESUME_LATEST:
        latest_ckpt = find_latest_checkpoint(CHECKPOINT_DIR)
        if latest_ckpt is not None:
            try:
                start_episode = ppo.load_checkpoint(latest_ckpt)
                print(
                    f"Resumed from checkpoint: {latest_ckpt} | "
                    f"saved_episode={start_episode} | training_step={ppo.training_step}"
                )
            except Exception as exc:  # pragma: no cover - runtime guard
                print(f"Failed to load checkpoint {latest_ckpt}: {exc}")
                print("Start training from scratch.")
                start_episode = 0
        else:
            print(f"No checkpoint found in {CHECKPOINT_DIR}, start from scratch.")

    try:
        for i_episode in range(start_episode, NUM_EPISODES):
            game.set_render_enabled((i_episode + 1) % RENDER_EVERY_EPISODES == 0)
            s, info = game.reset(seed=42 + i_episode)

            done = False
            ep_r = 0.0
            ep_steps = 0
            old_empty_cells = int(info.get("empty_cells", 0))
            old_max_block = int(info.get("max_block", np.max(s)))

            episode_states = []
            episode_actions = []
            episode_log_probs = []
            episode_valid_masks = []
            episode_rewards = []
            episode_dones = []
            episode_values = []

            while not done:
                valid_mask = ppo.get_valid_action_mask(s)
                a, log_prob, value = ppo.choose_action(s, valid_mask=valid_mask)
                s_, env_reward, done, info = game.step(a)

                max_block = int(info.get("max_block", np.max(s_)))
                is_success = bool(info.get("is_success", False))
                empty_cells = int(info.get("empty_cells", 0))

                r_merge = 0.3 * math.log1p(max(0.0, float(env_reward)))
                r_max = 0.5 * (math.log2(max_block) - math.log2(old_max_block))
                r_empty = 0.02 * (empty_cells - old_empty_cells)
                r_milestone = 0.0
                if max_block > old_max_block:
                    for tile_value, tile_bonus in MAX_TILE_MILESTONE_BONUSES:
                        if old_max_block < tile_value <= max_block:
                            r_milestone += tile_bonus
                invalid_move = np.array_equal(s_, s)
                r_invalid = -0.2 if invalid_move else 0.0
                r_done_fail = -0.8 if (done and not is_success) else 0.0
                r_success = 2.0 if is_success else 0.0

                r = (
                    r_merge
                    + r_max
                    + r_empty
                    + r_milestone
                    + r_invalid
                    + r_done_fail
                    + r_success
                )
                r = float(np.tanh(r / REWARD_TANH_SCALE) * REWARD_TANH_SCALE)

                episode_states.append(ppo._to_state_vector(s))
                episode_actions.append(a)
                episode_log_probs.append(log_prob)
                episode_valid_masks.append(np.asarray(valid_mask, dtype=np.bool_))
                episode_rewards.append(r)
                episode_dones.append(done)
                episode_values.append(value)

                ep_r += r
                ep_steps += 1
                timestep += 1

                s = s_
                old_empty_cells = empty_cells
                old_max_block = max_block

            advantages, returns = ppo.compute_gae(
                episode_rewards, episode_values, episode_dones, next_value=0.0
            )
            buffer["states"].extend(episode_states)
            buffer["actions"].extend(episode_actions)
            buffer["log_probs"].extend(episode_log_probs)
            buffer["valid_masks"].extend(episode_valid_masks)
            buffer["advantages"].extend(advantages.tolist())
            buffer["returns"].extend(returns.tolist())

            if timestep >= UPDATE_TIMESTEP:
                last_policy_loss, last_value_loss = ppo.learn(buffer)
                buffer = new_buffer()
                timestep = 0

            policy_text = (
                "None" if last_policy_loss is None else f"{last_policy_loss:.4f}"
            )
            value_text = "None" if last_value_loss is None else f"{last_value_loss:.4f}"
            print(
                f"Ep: {i_episode:4d} | Steps: {ep_steps:4d} | Ep_r: {ep_r:7.2f} | "
                f"PolicyLoss: {policy_text} | ValueLoss: {value_text} | Max_block: {old_max_block:5d}"
            )

            if (i_episode + 1) % SAVE_EVERY_EPISODES == 0:
                ckpt_path = CHECKPOINT_DIR / f"ppo_clip_ep_{i_episode + 1}.pt"
                ppo.save_checkpoint(ckpt_path, episode=i_episode + 1)
                print(f"Checkpoint saved: {ckpt_path}")

        # Final update for remaining samples.
        if len(buffer["states"]) > 0:
            ppo.learn(buffer)
    finally:
        game.close()

    print("Training completed!")


if __name__ == "__main__":
    main()
View Code

最后是一个加载模型进行推理的代码:

from __future__ import annotations

import argparse
from pathlib import Path

import torch

from ppo_clip import CHECKPOINT_DIR, PPO


def move_optimizer_state_to_device(optimizer: torch.optim.Optimizer, device: torch.device):
    """Ensure optimizer states are on the same device as model parameters."""
    for state in optimizer.state.values():
        for key, value in state.items():
            if torch.is_tensor(value):
                state[key] = value.to(device)


def load_checkpoint_for_resume(ckpt_path: Path) -> PPO:
    """Load full checkpoint and return a PPO agent ready to continue training."""
    agent = PPO()
    checkpoint = torch.load(ckpt_path, map_location=agent.device)

    agent.policy_net.load_state_dict(checkpoint["policy_net_state_dict"])
    agent.value_net.load_state_dict(checkpoint["value_net_state_dict"])
    agent.policy_optimizer.load_state_dict(checkpoint["policy_optimizer_state_dict"])
    agent.value_optimizer.load_state_dict(checkpoint["value_optimizer_state_dict"])
    agent.training_step = int(checkpoint.get("training_step", 0))

    move_optimizer_state_to_device(agent.policy_optimizer, agent.device)
    move_optimizer_state_to_device(agent.value_optimizer, agent.device)

    print(
        f"[Resume] Loaded: {ckpt_path}\n"
        f"  episode={checkpoint.get('episode', 'unknown')}, "
        f"training_step={agent.training_step}, device={agent.device}"
    )
    return agent


def load_policy_only(ckpt_path: Path) -> PPO:
    """Load only policy/value weights for inference, skip optimizer states."""
    agent = PPO()
    checkpoint = torch.load(ckpt_path, map_location=agent.device)

    agent.policy_net.load_state_dict(checkpoint["policy_net_state_dict"])
    agent.value_net.load_state_dict(checkpoint["value_net_state_dict"])
    agent.policy_net.eval()
    agent.value_net.eval()

    print(
        f"[Policy-Only] Loaded: {ckpt_path}\n"
        f"  episode={checkpoint.get('episode', 'unknown')}, device={agent.device}"
    )
    return agent


def select_action_greedy(agent: PPO, state, valid_mask) -> int:
    """Select action by max probability among valid actions."""
    state_vec = agent._to_state_vector(state)
    state_t = torch.as_tensor(
        state_vec, dtype=torch.float32, device=agent.device
    ).unsqueeze(0)
    mask_t = torch.as_tensor(valid_mask, dtype=torch.bool, device=agent.device)

    with torch.no_grad():
        probs = agent.policy_net(state_t).squeeze(0)
        probs = agent._apply_prob_mask(probs, mask_t)
        action = int(torch.argmax(probs).item())
    return action


def autoplay(
    agent: PPO,
    episodes: int,
    seed: int,
    fps: int,
    render: bool,
    deterministic: bool,
    max_steps: int,
):
    """Run automatic play loop by loading policy and calling Game2048 interfaces."""
    try:
        from model import Game2048, run
    except ModuleNotFoundError as exc:
        raise RuntimeError(
            "Autoplay requires gym dependency. Install in your conda env, e.g. "
            "`pip install gymnasium gym pygame`."
        ) from exc

    agent.policy_net.eval()
    agent.value_net.eval()

    game = Game2048(
        seed=seed,
        fps=fps,
        window_title="2048 PPO AutoPlay",
        render_enabled=render,
    )
    try:
        for ep in range(episodes):
            state, info = game.reset(seed=seed + ep)
            done = False
            step_count = 0
            env_reward_sum = 0.0

            while not done:
                valid_mask = agent.get_valid_action_mask(state)
                if deterministic:
                    action = select_action_greedy(agent, state, valid_mask)
                else:
                    action, _, _ = agent.choose_action(state, valid_mask=valid_mask)

                state, env_reward, done, info = run(game, action)
                env_reward_sum += float(env_reward)
                step_count += 1

                if max_steps > 0 and step_count >= max_steps:
                    done = True
                    break

                if info.get("closed_by_user", False):
                    done = True
                    break

            max_tile = int(info.get("max_block", info.get("max_tile", 0)))
            print(
                f"[AutoPlay] Episode {ep + 1}/{episodes} | "
                f"steps={step_count} | env_reward_sum={env_reward_sum:.2f} | max_tile={max_tile}"
            )

            if info.get("closed_by_user", False):
                print("Window closed by user, stop autoplay.")
                break
    finally:
        game.close()


def find_latest_checkpoint(ckpt_dir: Path) -> Path:
    """Find latest checkpoint by episode number in filename."""
    candidates = list(ckpt_dir.glob("ppo_clip_ep_*.pt"))
    if not candidates:
        raise FileNotFoundError(f"No checkpoint found in: {ckpt_dir}")

    def episode_from_name(path: Path) -> int:
        try:
            return int(path.stem.rsplit("_", 1)[-1])
        except (ValueError, IndexError):
            return -1

    return max(candidates, key=episode_from_name)


def parse_args():
    parser = argparse.ArgumentParser(description="Load PPO-Clip checkpoint demo.")
    parser.add_argument(
        "--ckpt",
        type=str,
        default=None,
        help="Checkpoint path. If omitted, auto-pick latest file in checkpoints dir.",
    )
    parser.add_argument(
        "--mode",
        type=str,
        default="autoplay",
        choices=["resume", "policy_only", "autoplay"],
        help=(
            "resume: load network + optimizer; "
            "policy_only: load weights only; "
            "autoplay: load policy and auto-play 2048."
        ),
    )
    parser.add_argument(
        "--episodes",
        type=int,
        default=3,
        help="How many episodes to autoplay in autoplay mode.",
    )
    parser.add_argument(
        "--seed",
        type=int,
        default=42,
        help="Base seed for autoplay mode.",
    )
    parser.add_argument(
        "--fps",
        type=int,
        default=30,
        help="Render FPS for autoplay mode.",
    )
    parser.add_argument(
        "--render",
        type=int,
        default=1,
        choices=[0, 1],
        help="Whether to render game window in autoplay mode: 1=yes, 0=no.",
    )
    parser.add_argument(
        "--deterministic",
        type=int,
        default=1,
        choices=[0, 1],
        help="Action selection in autoplay mode: 1=greedy, 0=sample.",
    )
    parser.add_argument(
        "--max_steps",
        type=int,
        default=0,
        help="Max steps per episode in autoplay mode. 0 means no limit.",
    )
    return parser.parse_args()


def main():
    args = parse_args()
    ckpt_path = Path(args.ckpt) if args.ckpt else find_latest_checkpoint(CHECKPOINT_DIR)

    if args.mode == "resume":
        _ = load_checkpoint_for_resume(ckpt_path)
        print("You can continue training using this returned PPO instance.")
    elif args.mode == "policy_only":
        _ = load_policy_only(ckpt_path)
        print("You can run inference/evaluation using this returned PPO instance.")
    else:
        agent = load_policy_only(ckpt_path)
        try:
            autoplay(
                agent=agent,
                episodes=max(1, int(args.episodes)),
                seed=int(args.seed),
                fps=max(1, int(args.fps)),
                render=bool(args.render),
                deterministic=bool(args.deterministic),
                max_steps=max(0, int(args.max_steps)),
            )
        except RuntimeError as exc:
            print(f"[AutoPlay Error] {exc}")


if __name__ == "__main__":
    main()
View Code

3.碎碎念

与此同时,带来了一些小问题,想多碎碎念几句

首先就是强化学习的训练

之前只是听学长说强化学习难训练,现在我是亲身体验了

训练真的很容易陷入局部最优解。像是我的2048,训练了2w轮,也只能停留在稳定512,再往上真的很难处理。

后面又修改了奖励函数和训练方法,也不是特别理想吧。

那这就会导致

训练成本可能比那些要卡的学习还要高

因为这个完全就是不可控的,谁也不知道策略会不会滑到某个神奇的角落

接下来就是奖励函数的设计

尽量要把奖励函数弄稠密些

这个倒是好办,把任务分解为小目标,然后对每个行为都规定一个小奖励和小惩罚就差不多了,到时候就具体问题具体分析了。

(end)

 

posted @ 2026-03-30 20:18  阿基米德的澡盆  阅读(63)  评论(0)    收藏  举报