53 lines
1.8 KiB
Python
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)
|
|
|