import torch import torch.nn as nn import torch.optim as optim from torch.distributions import Categorical from dlgo import encoders from dlgo import goboard from dlgo.agent import Agent from dlgo.agent.helpers import is_point_an_eye class ValueAgent(Agent): def __init__(self, model, encoder, policy='eps-greedy'): Agent.__init__(self) self.model = model self.encoder = encoder self.collector = None self.temperature = 0.0 self.policy = policy self.last_move_value = 0 def predict(self, game_state): encoded_state = self.encoder.encode(game_state) input_tensor = torch.tensor([encoded_state]) return self.model(input_tensor)[0] def set_temperature(self, temperature): self.temperature = temperature def set_collector(self, collector): self.collector = collector def set_policy(self, policy): if policy not in ('eps-greedy', 'weighted'): raise ValueError(policy) self.policy = policy def select_move(self, game_state): moves = [] board_tensors = [] for move in game_state.legal_moves(): if not move.is_play: continue next_state = game_state.apply_move(move) board_tensor = self.encoder.encode(next_state) moves.append(move) board_tensors.append(board_tensor) if not moves: return goboard.Move.pass_turn() board_tensors = torch.tensor(board_tensors) opp_values = self.model(board_tensors) opp_values = opp_values.reshape(len(moves)) values = 1 - opp_values if self.policy == 'eps-greedy': ranked_moves = self.rank_moves_eps_greedy(values) elif self.policy == 'weighted': ranked_moves = self.rank_moves_weighted(values) else: ranked_moves = None for move_idx in ranked_moves: move = moves[move_idx] if not is_point_an_eye(game_state.board, move.point, game_state.next_player): if self.collector is not None: self.collector.record_decision( state=board_tensor, action=self.encoder.encode_point(move.point), ) self.last_move_value = float(values[move_idx]) return move return goboard.Move.pass_turn() def rank_moves_eps_greedy(self, values): if torch.rand(1).item() < self.temperature: values = torch.rand_like(values) ranked_moves = torch.argsort(values) return ranked_moves[::-1] def rank_moves_weighted(self, values): p = values / torch.sum(values) p = torch.pow(p, 1.0 / self.temperature) p = p / torch.sum(p) return Categorical(p).sample().item() def train(self, experience, lr=0.1, batch_size=128): opt = optim.SGD(self.model.parameters(), lr=lr) criterion = nn.MSELoss() n = experience.states.shape[0] y = torch.zeros((n,)) for i in range(n): reward = experience.rewards[i] y[i] = 1 if reward > 0 else 0 self.model.train() opt.zero_grad() pred = self.model(experience.states) loss = criterion(pred, y) loss.backward() opt.step() def diagnostics(self): return {'value': self.last_move_value} def save(self, filename): torch.save({ 'model': self.model.state_dict(), 'encoder': { 'name': self.encoder.name(), 'board_width': self.encoder.board_width, 'board_height': self.encoder.board_height, } }, filename) @classmethod def load(cls, filename): checkpoint = torch.load(filename) model = Model() # The Model class should be defined appropriately model.load_state_dict(checkpoint['model']) encoder_info = checkpoint['encoder'] encoder = encoders.get_encoder_by_name( encoder_info['name'], (encoder_info['board_width'], encoder_info['board_height'])) return cls(model, encoder) def load_value_agent(filename): # 我假设 Model 类是预先定义好的 PyTorch 模型,你需要根据实际情况替换或实现这个模型。 checkpoint = torch.load(filename) model = Model() # The Model class should be defined appropriately model.load_state_dict(checkpoint['model']) encoder_info = checkpoint['encoder'] encoder = encoders.get_encoder_by_name( encoder_info['name'], (encoder_info['board_width'], encoder_info['board_height'])) return ValueAgent(model, encoder)