1
This commit is contained in:
@@ -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()
|
||||
Reference in New Issue
Block a user