mirror of
https://github.com/wassname/Castor.git
synced 2026-09-09 11:13:20 +08:00
Add hyperparameter tuning script
This commit is contained in:
+7
-5
@@ -51,8 +51,12 @@ def create_context(config):
|
||||
test_loader = utils.data.DataLoader(test_set, batch_size=1, collate_fn=collate_fn)
|
||||
|
||||
params = list(filter(lambda x: x.requires_grad, model.parameters()))
|
||||
optimizer = optim.Adam(params, lr=config.lr, weight_decay=config.weight_decay)
|
||||
# optimizer = optim.SGD(params, lr=config.lr, momentum=config.momentum, weight_decay=config.weight_decay)
|
||||
if config.optimizer == "adam":
|
||||
optimizer = optim.Adam(params, lr=config.lr, weight_decay=config.weight_decay)
|
||||
elif config.optimizer == "sgd":
|
||||
optimizer = optim.SGD(params, lr=config.lr, momentum=config.momentum, weight_decay=config.weight_decay)
|
||||
elif config.optimizer == "rmsprop":
|
||||
optimizer = optim.RMSprop(params, lr=config.lr, alpha=config.decay, momentum=config.momentum, weight_decay=config.weight_decay)
|
||||
criterion = nn.KLDivLoss()
|
||||
log_writer = LogWriter()
|
||||
return Context(model, train_loader, dev_loader, test_loader, optimizer, criterion, params, log_writer)
|
||||
@@ -95,7 +99,7 @@ def train(config):
|
||||
if i % config.mbatch_size == (config.mbatch_size - 1):
|
||||
loss /= config.mbatch_size
|
||||
loss.backward()
|
||||
nn.utils.clip_grad_norm(context.params, 5)
|
||||
nn.utils.clip_grad_norm(context.params, config.clip_norm)
|
||||
context.optimizer.step()
|
||||
|
||||
loss = loss.cpu().data[0]
|
||||
@@ -108,8 +112,6 @@ def train(config):
|
||||
best_dev_pr = result.pearsonr
|
||||
print("Saving best model...")
|
||||
context.model.save(config.output_file)
|
||||
test_result = evaluate(context, context.test_loader)
|
||||
print("Final test result: {}".format(test_result))
|
||||
|
||||
def main():
|
||||
config = data.Configs.base_config()
|
||||
|
||||
+5
-1
@@ -9,6 +9,8 @@ class Configs(object):
|
||||
@staticmethod
|
||||
def base_config():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--classifier", type=str, default="vdpwi", choices=["vdpwi", "resnet"])
|
||||
parser.add_argument("--clip_norm", type=float, default=5)
|
||||
parser.add_argument("--cpu", action="store_true", default=False)
|
||||
parser.add_argument("--dataset", type=str, default="sick", choices=["sick"])
|
||||
parser.add_argument("--decay", type=float, default=0.95)
|
||||
@@ -17,10 +19,12 @@ class Configs(object):
|
||||
parser.add_argument("--mbatch_size", type=int, default=16)
|
||||
parser.add_argument("--mode", type=str, default="train", choices=["train", "test"])
|
||||
parser.add_argument("--momentum", type=float, default=0.9)
|
||||
parser.add_argument("--n_epochs", type=int, default=40)
|
||||
parser.add_argument("--n_epochs", type=int, default=35)
|
||||
parser.add_argument("--n_labels", type=int, default=5)
|
||||
parser.add_argument("--optimizer", type=str, default="adam", choices=["adam", "sgd", "rmsprop"])
|
||||
parser.add_argument("--output_file", type=str, default="local_saves/model.pt")
|
||||
parser.add_argument("--res_fmaps", type=int, default=32)
|
||||
parser.add_argument("--res_layers", type=int, default=16)
|
||||
parser.add_argument("--restore", action="store_true", default=False)
|
||||
parser.add_argument("--rnn_hidden_dim", type=int, default=250)
|
||||
parser.add_argument("--weight_decay", type=float, default=5E-4)
|
||||
|
||||
+42
-12
@@ -2,6 +2,7 @@ from torch.autograd import Variable
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torchvision.models as models
|
||||
import numpy as np
|
||||
|
||||
class SerializableModule(nn.Module):
|
||||
@@ -14,8 +15,41 @@ class SerializableModule(nn.Module):
|
||||
def load(self, filename):
|
||||
self.load_state_dict(torch.load(filename, map_location=lambda storage, loc: storage))
|
||||
|
||||
def hard_pad2d(x, pad):
|
||||
def pad_side(idx):
|
||||
pad_len = max(pad - x.size(idx), 0)
|
||||
return [0, pad_len]
|
||||
padding = pad_side(3)
|
||||
padding.extend(pad_side(2))
|
||||
x = F.pad(x, padding)
|
||||
return x[:, :, :pad, :pad]
|
||||
|
||||
class ResNet(SerializableModule):
|
||||
def __init__(self, config):
|
||||
super().__init__()
|
||||
n_layers = config.res_layers
|
||||
n_maps = config.res_fmaps
|
||||
n_labels = config.n_labels
|
||||
self.conv0 = nn.Conv2d(12, n_maps, (3, 3), padding=1)
|
||||
self.convs = [nn.Conv2d(n_maps, n_maps, (3, 3), padding=1) for _ in range(n_layers)]
|
||||
self.output = nn.Linear(n_maps, n_labels)
|
||||
self.input_len = None
|
||||
for i, conv in enumerate(self.convs):
|
||||
self.add_module("conv{}".format(i + 1), conv)
|
||||
|
||||
def forward(self, x):
|
||||
x = F.relu(self.conv0(x))
|
||||
old_x = x
|
||||
for i, conv in enumerate(self.convs):
|
||||
x = F.relu(conv(x))
|
||||
if i % 2 == 1:
|
||||
x += old_x
|
||||
old_x = x
|
||||
x = torch.mean(x.view(x.size(0), x.size(1), -1), 2)
|
||||
return self.output(x)
|
||||
|
||||
class VDPWIConvNet(SerializableModule):
|
||||
def __init__(self, n_labels):
|
||||
def __init__(self, config):
|
||||
super().__init__()
|
||||
self.conv1 = nn.Conv2d(12, 128, 3, padding=1)
|
||||
self.conv2 = nn.Conv2d(128, 164, 3, padding=1)
|
||||
@@ -24,18 +58,11 @@ class VDPWIConvNet(SerializableModule):
|
||||
self.conv5 = nn.Conv2d(192, 128, 3, padding=1)
|
||||
self.maxpool2 = nn.MaxPool2d(2, ceil_mode=True)
|
||||
self.dnn = nn.Linear(128, 128)
|
||||
self.output = nn.Linear(128, n_labels)
|
||||
self.output = nn.Linear(128, config.n_labels)
|
||||
self.input_len = 32
|
||||
|
||||
def forward(self, x):
|
||||
def pad_side(idx):
|
||||
pad_len = max(32 - x.size(idx), 0)
|
||||
return [0, pad_len]
|
||||
padding = pad_side(3)
|
||||
padding.extend(pad_side(2))
|
||||
x = F.pad(x, padding)
|
||||
x = x[:, :, :32, :32]
|
||||
|
||||
x = hard_pad2d(x, self.input_len)
|
||||
pool_final = nn.MaxPool2d(2, ceil_mode=True) if x.size(2) == 32 else nn.MaxPool2d(3, 1, ceil_mode=True)
|
||||
x = self.maxpool2(F.relu(self.conv1(x)))
|
||||
x = self.maxpool2(F.relu(self.conv2(x)))
|
||||
@@ -46,13 +73,16 @@ class VDPWIConvNet(SerializableModule):
|
||||
return self.output(x)
|
||||
|
||||
class VDPWIModel(SerializableModule):
|
||||
def __init__(self, embedding, config, classifier_net=None):
|
||||
def __init__(self, embedding, config):
|
||||
super().__init__()
|
||||
self.hidden_dim = config.rnn_hidden_dim
|
||||
self.rnn = nn.LSTM(300, self.hidden_dim, 1, batch_first=True)
|
||||
self.embedding = embedding
|
||||
self.use_cuda = not config.cpu
|
||||
self.classifier_net = VDPWIConvNet(config.n_labels) if classifier_net is None else classifier_net
|
||||
if config.classifier == "vdpwi":
|
||||
self.classifier_net = VDPWIConvNet(config)
|
||||
elif config.classifier == "resnet":
|
||||
self.classifier_net = ResNet(config)
|
||||
|
||||
def compute_sim_cube(self, seq1, seq2, truncate=None):
|
||||
def compute_sim(prism1, prism2):
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
import os
|
||||
import random
|
||||
|
||||
class RandomParamIterator(object):
|
||||
def __init__(self, param_sets):
|
||||
self.param_sets = param_sets
|
||||
|
||||
def random_param_set(self):
|
||||
param_set = {}
|
||||
for param_key, param_values in self.param_sets.items():
|
||||
param_set[param_key] = random.choice(param_values)
|
||||
return param_set
|
||||
|
||||
class Tuner(object):
|
||||
def __init__(self, *iterators, limit=100):
|
||||
self.iterators = iterators
|
||||
self.limit = limit
|
||||
|
||||
def start(self):
|
||||
for i in range(self.limit):
|
||||
iterator = random.choice(self.iterators)
|
||||
params = iterator.random_param_set()
|
||||
print(params)
|
||||
arg_str = " ".join("--{}={}".format(k, v) for k, v in params.items())
|
||||
os.system("python . {} --output_file local_saves/model{}.pt".format(arg_str, i))
|
||||
|
||||
def main():
|
||||
vgg_param_sets = dict(classifer=["vdpwi"], clip_norm=[3, 5, 7], decay=[0.9, 0.95], lr=[5E-3, 1E-3, 5E-4],
|
||||
mbatch_size=[8, 16, 32], optimizer=["adam", "rmsprop"], rnn_hidden_dim=[150, 250, 300],
|
||||
weight_decay=[0, 5E-4, 1E-3])
|
||||
res_param_sets = dict(classifier=["resnet"], clip_norm=[3, 5, 7], decay=[0.9, 0.95], lr=[5E-3, 1E-3, 5E-4],
|
||||
mbatch_size=[8, 16, 32], rnn_hidden_dim=[150, 250, 300], res_fmaps=[16, 24, 32], res_layers=[4, 8, 16, 24])
|
||||
vgg_iterator = RandomParamIterator(vgg_param_sets)
|
||||
res_iterator = RandomParamIterator(res_param_sets)
|
||||
Tuner(vgg_iterator, res_iterator).start()
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user