""" supervised learning,采用监督学习的策略网络 """ from tugo.data_processing.parallel_processor import GoDataProcessor from tugo.encoders.alphago import AlphaGoEncoder from tugo.agents.predict import DeepLearningAgent from tugo.models.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()