Files
tugo/rl/value.py
T
2023-05-30 18:48:27 +08:00

150 lines
4.7 KiB
Python

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)