update save and load

This commit is contained in:
2023-05-31 16:16:24 +08:00
parent 41c70802fd
commit fe2c327f4e
12 changed files with 252 additions and 44 deletions
+111
View File
@@ -0,0 +1,111 @@
import torch
import numpy as np
from tugo import encoders
from tugo.game_logic import goboard
from tugo.agents import Agent
class ACAgent(Agent):
def __init__(self, model, encoder):
Agent.__init__(self)
self.model = model
self.encoder = encoder
self.collector = None
self.temperature = 1.0
self.last_state_value = 0
def set_temperature(self, temperature):
self.temperature = temperature
def set_collector(self, collector):
self.collector = collector
def select_move(self, game_state):
num_moves = self.encoder.board_width * self.encoder.board_height
board_tensor = self.encoder.encode(game_state)
x = torch.from_numpy(np.array([board_tensor])).float().to(self.model.device)
actions, values = self.model(x)
move_probs = actions[0].detach().cpu().numpy()
estimated_value = values[0][0].item()
self.last_state_value = float(estimated_value)
# Prevent move probs from getting stuck at 0 or 1.
move_probs = np.power(move_probs, 1.0 / self.temperature)
move_probs = move_probs / np.sum(move_probs)
eps = 1e-6
move_probs = np.clip(move_probs, eps, 1 - eps)
# Re-normalize to get another probability distribution.
move_probs = move_probs / np.sum(move_probs)
# Turn the probabilities into a ranked list of moves.
candidates = np.arange(num_moves)
ranked_moves = np.random.choice(
candidates, num_moves, replace=False, p=move_probs)
for point_idx in ranked_moves:
point = self.encoder.decode_point_index(point_idx)
true_move = goboard.Move.play(point)
if not game_state.is_valid_move(true_move):
true_move = goboard.Move.pass_turn()
if self.collector is not None:
self.collector.record_decision(
state=board_tensor,
action=point_idx,
estimated_value=estimated_value
)
return true_move
# No legal, non-self-destructive moves less.
return goboard.Move.pass_turn()
def train(self, experience, lr=0.1, batch_size=128):
optimizer = torch.optim.SGD(self.model.parameters(), lr=lr, clipvalue=0.2)
criterion = torch.nn.CrossEntropyLoss()
n = experience.states.shape[0]
num_moves = self.encoder.num_points()
policy_target = np.zeros((n, num_moves))
value_target = np.zeros((n,))
for i in range(n):
action = experience.actions[i]
reward = experience.rewards[i]
policy_target[i][action] = experience.advantages[i]
value_target[i] = reward
self.model.train()
optimizer.zero_grad()
policy_target_t = torch.from_numpy(policy_target).float().to(self.model.device)
value_target_t = torch.from_numpy(value_target).float().to(self.model.device)
states_t = torch.from_numpy(experience.states).float().to(self.model.device)
actions_pred, values_pred = self.model(states_t)
loss1 = criterion(actions_pred, policy_target_t)
loss2 = criterion(values_pred, value_target_t)
loss = loss1 + loss2
loss.backward()
optimizer.step()
def serialize(self, file_path):
torch.save({
'model_state_dict': self.model.state_dict(),
'encoder_name': self.encoder.name(),
'encoder_board_width': self.encoder.board_width,
'encoder_board_height': self.encoder.board_height,
}, file_path)
def diagnostics(self):
return {'value': self.last_state_value}
def load_passing_ac_agent(file_path):
checkpoint = torch.load(file_path)
model = TheModelClass() # TheModelClass should be replaced with your model class
model.load_state_dict(checkpoint['model_state_dict'])
encoder_name = checkpoint['encoder_name']
board_width = checkpoint['encoder_board_width']
board_height = checkpoint['encoder_board_height']
encoder = encoders.get_encoder_by_name(
encoder_name,
(board_width, board_height))
return ACAgent(model, encoder)