Files
2023-06-01 17:23:39 +08:00

151 lines
4.8 KiB
Python

import torch
import torch.nn as nn
import torch.optim as optim
from torch.distributions import Categorical
from tugo import encoders
# from tugo.game_logic import goboard
from tugo.game_logic import goboard_fast as goboard
from tugo.agents import Agent
from tugo.agents.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, file_path):
torch.save({
'model': self.model.state_dict(),
'encoder': {
'name': self.encoder.name(),
'board_width': self.encoder.board_width,
'board_height': self.encoder.board_height,
}
}, file_path)
@classmethod
def load(cls, file_path):
checkpoint = torch.load(file_path)
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(file_path):
# 我假设 Model 类是预先定义好的 PyTorch 模型,你需要根据实际情况替换或实现这个模型。
checkpoint = torch.load(file_path)
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)