96 lines
2.9 KiB
Python
96 lines
2.9 KiB
Python
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()
|