Files
tugo/models/alphago.py
T
2023-05-30 17:16:48 +08:00

53 lines
1.8 KiB
Python

import torch
import torch.nn as nn
class AlphaGoModel(nn.Module):
def __init__(self, input_shape, is_policy_net=False,
# num_filters=192, first_kernel_size=5, other_kernel_size=3):
num_filters=192, first_kernel_size=5, other_kernel_size=3, num_classes=None):
super(AlphaGoModel, self).__init__()
layers = [
nn.Conv2d(input_shape[0], num_filters, first_kernel_size, padding=first_kernel_size//2),
nn.ReLU()
]
for i in range(2, 12):
layers.extend([
nn.Conv2d(num_filters, num_filters, other_kernel_size, padding=other_kernel_size//2),
nn.ReLU()
])
if is_policy_net:
assert num_classes is not None, "num_classes must be provided for policy network"
layers.extend([
nn.Conv2d(num_filters, 1, kernel_size=1, padding=0),
# nn.Softmax(dim=1),
nn.Flatten()
])
else:
layers.extend([
nn.Conv2d(num_filters, num_filters, other_kernel_size, padding=other_kernel_size//2),
nn.ReLU(),
nn.Conv2d(num_filters, 1, kernel_size=1, padding=0),
nn.ReLU(),
nn.Flatten(),
nn.Linear(input_shape[1] * input_shape[2], 256),
nn.ReLU(),
nn.Linear(256, 1),
nn.Tanh()
])
self.model = nn.Sequential(*layers)
def forward(self, x):
return self.model(x)
# Instantiate policy network
# alphago_policy_model = AlphaGoModel(input_shape=(3, 19, 19), is_policy_net=True)
alphago_policy_model = AlphaGoModel(input_shape=(3, 19, 19), is_policy_net=True, num_classes=361)
# Instantiate value network
alphago_value_model = AlphaGoModel(input_shape=(3, 19, 19), is_policy_net=False)