This commit is contained in:
2023-05-23 15:52:09 +08:00
parent e8f9c8a287
commit 663edbb5e2
88 changed files with 6250 additions and 0 deletions
+95
View File
@@ -0,0 +1,95 @@
from tugo.data.parallel_processor import GoDataProcessor
from tugo.encoders.alphago import AlphaGoEncoder
from tugo.agent.predict import DeepLearningAgent
from tugo.networks.alphago import AlphaGoModel
import torch
from torch import nn
import torch.optim as optim
from torch.utils.data import DataLoader
from torch.utils.tensorboard import SummaryWriter
from tqdm import tqdm
rows, cols = 19, 19
num_classes = rows * cols
# num_games = 10000
num_games = 1000
encoder = AlphaGoEncoder()
processor = GoDataProcessor(encoder=encoder.name())
train_dataset = processor.load_go_data('train', num_games, use_generator=False)
test_dataset = processor.load_go_data('test', num_games, use_generator=False)
input_shape = (encoder.num_planes, rows, cols)
# alphago_sl_policy = AlphaGoModel(input_shape, is_policy_net=True)
alphago_sl_policy = AlphaGoModel(input_shape, is_policy_net=True, num_classes=361)
optimizer = optim.SGD(alphago_sl_policy.parameters(), lr=0.01)
criterion = nn.CrossEntropyLoss()
# Set up TensorBoard logging
summary_writer = SummaryWriter(log_dir='./logs')
epochs = 200
# batch_size = 128
batch_size = 1024
# epochs = 1
# batch_size = 32
train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)
test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False)
# 检查是否有可用的 GPU
if torch.cuda.is_available():
print("cuda is available, will use gpu!")
device = torch.device('cuda')
else:
print("cuda is unavailable, will use cpu")
device = torch.device('cpu')
alphago_sl_policy.to(device)
alphago_sl_policy.train()
for epoch in range(epochs):
print(f"Epoch: {epoch+1}/{epochs}")
for step, (inputs, targets) in enumerate(tqdm(train_loader)):
# print("Inputs shape:", inputs.shape)
# print("Targets shape:", targets.shape)
# print("step:", step)
inputs, targets = inputs.to(device), targets.to(device)
optimizer.zero_grad()
outputs = alphago_sl_policy(inputs)
loss = criterion(outputs, targets)
loss.backward()
optimizer.step()
summary_writer.add_scalar('Training Loss', loss.item(), epoch * len(train_loader) + step)
alphago_sl_agent = DeepLearningAgent(alphago_sl_policy, encoder)
alphago_sl_agent.save('checkpoints/alphago_sl_policy.pt')
# Test set evaluation
# 评估测试集上的模型性能
alphago_sl_policy.eval()
test_loss = 0.0
correct = 0
total = 0
with torch.no_grad():
for inputs, targets in test_loader:
inputs, targets = inputs.to(device), targets.to(device)
outputs = alphago_sl_policy(inputs)
loss = criterion(outputs, targets)
test_loss += loss.item()
_, predicted = outputs.max(1)
total += targets.size(0)
correct += predicted.eq(targets).sum().item()
print(f"Test Loss: {test_loss/total:.4f}")
print(f"Test Accuracy: {correct/total:.4f}")
summary_writer.close()