mirror of
https://github.com/wassname/Castor.git
synced 2026-09-09 11:13:20 +08:00
Make model parallel wrt batch size
This commit is contained in:
+32
-17
@@ -22,22 +22,40 @@ def create_context(config):
|
||||
emb2 = []
|
||||
labels = []
|
||||
cmp_labels = []
|
||||
pad_cube = []
|
||||
max_len1 = 0; max_len2 = 0
|
||||
|
||||
for s1, s2, l, cl in batch:
|
||||
emb1.append(s1)
|
||||
emb2.append(s2)
|
||||
max_len1 = max(max_len1, len(s1))
|
||||
max_len2 = max(max_len2, len(s2))
|
||||
labels.append(l)
|
||||
cmp_labels.append(cl)
|
||||
|
||||
for s1, s2 in zip(emb1, emb2):
|
||||
pad1 = (max_len1 - len(s1))
|
||||
pad2 = (max_len2 - len(s2))
|
||||
pad_mask = np.ones((max_len1, max_len2))
|
||||
pad_mask[:len(s1), :len(s2)] = 0
|
||||
pad_cube.append(pad_mask)
|
||||
s1.extend([embedding.weight.size(0) - 1] * pad1)
|
||||
s2.extend([embedding.weight.size(0) - 1] * pad2)
|
||||
|
||||
pad_cube = np.array(pad_cube)
|
||||
emb1 = torch.LongTensor(emb1)
|
||||
emb2 = torch.LongTensor(emb2)
|
||||
labels = torch.Tensor(labels)
|
||||
emb1 = torch.autograd.Variable(emb1, requires_grad=False)
|
||||
emb2 = torch.autograd.Variable(emb2, requires_grad=False)
|
||||
labels = torch.autograd.Variable(labels, requires_grad=False)
|
||||
pad_cube = torch.autograd.Variable(torch.from_numpy(pad_cube).float(), requires_grad=False)
|
||||
if not config.cpu:
|
||||
emb1 = emb1.cuda()
|
||||
emb2 = emb2.cuda()
|
||||
labels = labels.cuda()
|
||||
return emb1, emb2, labels, cmp_labels
|
||||
pad_cube = pad_cube.cuda()
|
||||
return emb1, emb2, labels, pad_cube, cmp_labels
|
||||
|
||||
embedding, (train_set, dev_set, test_set) = data.load_dataset(config.dataset)
|
||||
model = mod.VDPWIModel(embedding, config)
|
||||
@@ -46,7 +64,7 @@ def create_context(config):
|
||||
if not config.cpu:
|
||||
model = model.cuda()
|
||||
|
||||
train_loader = utils.data.DataLoader(train_set, shuffle=True, batch_size=1, collate_fn=collate_fn)
|
||||
train_loader = utils.data.DataLoader(train_set, shuffle=True, batch_size=config.mbatch_size, collate_fn=collate_fn)
|
||||
dev_loader = utils.data.DataLoader(dev_set, batch_size=1, collate_fn=collate_fn)
|
||||
test_loader = utils.data.DataLoader(test_set, batch_size=1, collate_fn=collate_fn)
|
||||
|
||||
@@ -71,8 +89,8 @@ def evaluate(context, data_loader):
|
||||
model.eval()
|
||||
predictions = []
|
||||
true_labels = []
|
||||
for sent1, sent2, _, truth in data_loader:
|
||||
scores = model(sent1, sent2)
|
||||
for sent1, sent2, _, pad_cube, truth in data_loader:
|
||||
scores = model(sent1, sent2, pad_cube)
|
||||
scores = F.softmax(scores).cpu().data.numpy()[0]
|
||||
prediction = np.dot(np.arange(1, len(scores) + 1), scores)
|
||||
predictions.append(prediction); true_labels.append(truth[0][0])
|
||||
@@ -88,24 +106,21 @@ def train(config):
|
||||
best_dev_pr = 0
|
||||
for epoch_no in range(config.n_epochs):
|
||||
print("Epoch number: {}".format(epoch_no + 1))
|
||||
loader_wrapper = tqdm(enumerate(context.train_loader), total=len(context.train_loader), desc="Loss")
|
||||
loader_wrapper = tqdm(context.train_loader, total=len(context.train_loader), desc="Loss")
|
||||
context.model.train()
|
||||
loss = 0
|
||||
for i, (sent1, sent2, label_pmf, _) in loader_wrapper:
|
||||
for sent1, sent2, label_pmf, pad_cube, _ in loader_wrapper:
|
||||
context.optimizer.zero_grad()
|
||||
scores = F.log_softmax(context.model(sent1, sent2))
|
||||
scores = F.log_softmax(context.model(sent1, sent2, pad_cube))
|
||||
|
||||
loss = context.criterion(scores, label_pmf) + loss
|
||||
if i % config.mbatch_size == (config.mbatch_size - 1):
|
||||
loss /= config.mbatch_size
|
||||
loss.backward()
|
||||
nn.utils.clip_grad_norm(context.params, config.clip_norm)
|
||||
context.optimizer.step()
|
||||
loss = context.criterion(scores, label_pmf)
|
||||
loss.backward()
|
||||
nn.utils.clip_grad_norm(context.params, config.clip_norm)
|
||||
context.optimizer.step()
|
||||
|
||||
loss = loss.cpu().data[0]
|
||||
loader_wrapper.set_description("Loss: {:<8}".format(round(loss, 5)))
|
||||
context.log_writer.log_train_loss(loss)
|
||||
loss = 0
|
||||
loss = loss.cpu().data[0]
|
||||
loader_wrapper.set_description("Loss: {:<8}".format(round(loss, 5)))
|
||||
context.log_writer.log_train_loss(loss)
|
||||
result = evaluate(context, context.dev_loader)
|
||||
print("Dev result: {}".format(result))
|
||||
if best_dev_pr < result.pearsonr:
|
||||
|
||||
+5
-5
@@ -10,24 +10,24 @@ class Configs(object):
|
||||
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("--clip_norm", type=float, default=50)
|
||||
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)
|
||||
parser.add_argument("--input_file", type=str, default="local_saves/model.pt")
|
||||
parser.add_argument("--lr", type=float, default=1E-3)
|
||||
parser.add_argument("--lr", type=float, default=5E-4)
|
||||
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("--momentum", type=float, default=0.1)
|
||||
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("--optimizer", type=str, default="rmsprop", 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)
|
||||
parser.add_argument("--weight_decay", type=float, default=1E-5)
|
||||
parser.add_argument("--wordvecs_file", type=str, default="local_data/glove/glove.840B.300d.txt")
|
||||
return parser.parse_known_args()[0]
|
||||
|
||||
|
||||
+60
-46
@@ -51,11 +51,16 @@ class ResNet(SerializableModule):
|
||||
class VDPWIConvNet(SerializableModule):
|
||||
def __init__(self, config):
|
||||
super().__init__()
|
||||
self.conv1 = nn.Conv2d(12, 128, 3, padding=1)
|
||||
self.conv2 = nn.Conv2d(128, 164, 3, padding=1)
|
||||
self.conv3 = nn.Conv2d(164, 192, 3, padding=1)
|
||||
self.conv4 = nn.Conv2d(192, 192, 3, padding=1)
|
||||
self.conv5 = nn.Conv2d(192, 128, 3, padding=1)
|
||||
def make_conv(n_in, n_out):
|
||||
conv = nn.Conv2d(n_in, n_out, 3, padding=1)
|
||||
conv.bias.data.zero_()
|
||||
nn.init.xavier_normal(conv.weight)
|
||||
return conv
|
||||
self.conv1 = make_conv(12, 128)
|
||||
self.conv2 = make_conv(128, 164)
|
||||
self.conv3 = make_conv(164, 192)
|
||||
self.conv4 = make_conv(192, 192)
|
||||
self.conv5 = make_conv(192, 128)
|
||||
self.maxpool2 = nn.MaxPool2d(2, ceil_mode=True)
|
||||
self.dnn = nn.Linear(128, 128)
|
||||
self.output = nn.Linear(128, config.n_labels)
|
||||
@@ -84,60 +89,67 @@ class VDPWIModel(SerializableModule):
|
||||
elif config.classifier == "resnet":
|
||||
self.classifier_net = ResNet(config)
|
||||
|
||||
def compute_sim_cube(self, seq1, seq2, truncate=None):
|
||||
def compute_sim_cube(self, seq1, seq2):
|
||||
def compute_sim(prism1, prism2):
|
||||
prism1_len = prism1.norm(dim=2)
|
||||
prism2_len = prism2.norm(dim=2)
|
||||
prism1_len = prism1.norm(dim=3)
|
||||
prism2_len = prism2.norm(dim=3)
|
||||
|
||||
dot_prod = torch.matmul(prism1.unsqueeze(2), prism2.unsqueeze(3))
|
||||
dot_prod = dot_prod.squeeze(2).squeeze(2)
|
||||
dot_prod = torch.matmul(prism1.unsqueeze(3), prism2.unsqueeze(4))
|
||||
dot_prod = dot_prod.squeeze(3).squeeze(3)
|
||||
cos_dist = dot_prod / (prism1_len * prism2_len + 1E-8)
|
||||
l2_dist = (prism1 - prism2).norm(dim=2)
|
||||
return torch.stack([dot_prod, cos_dist, l2_dist], 0)
|
||||
l2_dist = -((prism1 - prism2).norm(dim=3))
|
||||
return torch.stack([dot_prod, cos_dist, l2_dist], 1)
|
||||
|
||||
def compute_prism(seq1, seq2):
|
||||
prism1 = seq1.repeat(seq2.size(0), 1, 1)
|
||||
prism2 = seq2.repeat(seq1.size(0), 1, 1)
|
||||
prism1 = prism1.permute(1, 0, 2).contiguous()
|
||||
prism2 = prism2.permute(0, 1, 2).contiguous()
|
||||
prism1 = seq1.repeat(seq2.size(1), 1, 1, 1)
|
||||
prism2 = seq2.repeat(seq1.size(1), 1, 1, 1)
|
||||
prism1 = prism1.permute(1, 2, 0, 3).contiguous()
|
||||
prism2 = prism2.permute(1, 0, 2, 3).contiguous()
|
||||
return compute_sim(prism1, prism2)
|
||||
|
||||
sim_cube = Variable(torch.Tensor(12, seq1.size(0), seq2.size(0)))
|
||||
sim_cube = Variable(torch.Tensor(seq1.size(0), 12, seq1.size(1), seq2.size(1)))
|
||||
if self.use_cuda:
|
||||
sim_cube = sim_cube.cuda()
|
||||
seq1_f = seq1[:, :self.hidden_dim]
|
||||
seq1_b = seq1[:, self.hidden_dim:]
|
||||
seq2_f = seq2[:, :self.hidden_dim]
|
||||
seq2_b = seq2[:, self.hidden_dim:]
|
||||
sim_cube[0:3] = compute_prism(seq1, seq2)
|
||||
sim_cube[3:6] = compute_prism(seq1_f, seq2_f)
|
||||
sim_cube[6:9] = compute_prism(seq1_b, seq2_b)
|
||||
sim_cube[9:12] = compute_prism(seq1_f + seq1_b, seq2_f + seq2_b)
|
||||
if truncate is not None:
|
||||
sim_cube = sim_cube[:, :truncate, :truncate].contiguous()
|
||||
seq1_f = seq1[:, :, :self.hidden_dim]
|
||||
seq1_b = seq1[:, :, self.hidden_dim:]
|
||||
seq2_f = seq2[:, :, :self.hidden_dim]
|
||||
seq2_b = seq2[:, :, self.hidden_dim:]
|
||||
sim_cube[:, 0:3] = compute_prism(seq1, seq2)
|
||||
sim_cube[:, 3:6] = compute_prism(seq1_f, seq2_f)
|
||||
sim_cube[:, 6:9] = compute_prism(seq1_b, seq2_b)
|
||||
sim_cube[:, 9:12] = compute_prism(seq1_f + seq1_b, seq2_f + seq2_b)
|
||||
return sim_cube
|
||||
|
||||
def compute_focus_cube(self, sim_cube):
|
||||
def compute_focus_cube(self, sim_cube, pad_cube):
|
||||
neg_magic = -10000
|
||||
pad_cube = pad_cube.repeat(12, 1, 1, 1)
|
||||
pad_cube = pad_cube.permute(1, 0, 2, 3).contiguous()
|
||||
sim_cube = neg_magic * pad_cube + sim_cube
|
||||
mask = Variable(torch.Tensor(*sim_cube.size()))
|
||||
if self.use_cuda:
|
||||
mask = mask.cuda()
|
||||
mask[:, :, :] = 0.1
|
||||
mask[:, :, :, :] = 0.1
|
||||
|
||||
def build_mask(index):
|
||||
s1tag = np.zeros(sim_cube.size(1))
|
||||
s2tag = np.zeros(sim_cube.size(2))
|
||||
_, indices = torch.sort(sim_cube[index].view(-1), descending=True)
|
||||
for i, index in enumerate(indices.cpu().data.numpy()):
|
||||
if i >= len(s1tag) + len(s2tag):
|
||||
break
|
||||
pos1, pos2 = index // len(s2tag), index % len(s2tag)
|
||||
if s1tag[pos1] + s2tag[pos2] == 0:
|
||||
s1tag[pos1] = s2tag[pos2] = 1
|
||||
mask[:, int(pos1), int(pos2)] = 1
|
||||
max_mask = sim_cube[:, index].clone()
|
||||
for _ in range(min(sim_cube.size(2), sim_cube.size(3))):
|
||||
values, indices = torch.max(max_mask.view(sim_cube.size(0), -1), 1)
|
||||
row_indices = indices / sim_cube.size(3)
|
||||
col_indices = indices % sim_cube.size(3)
|
||||
row_indices = row_indices.unsqueeze(1)
|
||||
col_indices = col_indices.unsqueeze(1).unsqueeze(1)
|
||||
for i, (row_i, col_i, val) in enumerate(zip(row_indices, col_indices, values)):
|
||||
if val < neg_magic / 2:
|
||||
continue
|
||||
mask[i, :, row_i, col_i] = 1
|
||||
max_mask[i, row_i, :] = neg_magic
|
||||
max_mask[i, :, col_i] = neg_magic
|
||||
build_mask(9)
|
||||
build_mask(10)
|
||||
return mask * sim_cube
|
||||
focus_cube = mask * sim_cube * (1 - pad_cube)
|
||||
return focus_cube
|
||||
|
||||
def forward(self, x1, x2):
|
||||
def forward(self, x1, x2, pad_cube):
|
||||
x1 = self.embedding(x1)
|
||||
x2 = self.embedding(x2)
|
||||
seq1f, _ = self.rnn(x1)
|
||||
@@ -146,9 +158,11 @@ class VDPWIModel(SerializableModule):
|
||||
seq2b, _ = self.rnn(torch.cat(x2.split(1, 1)[::-1], 1))
|
||||
seq1 = torch.cat([seq1f, seq1b], 2)
|
||||
seq2 = torch.cat([seq2f, seq2b], 2)
|
||||
seq1 = seq1.squeeze(0) # batch size assumed to be 1
|
||||
seq2 = seq2.squeeze(0)
|
||||
sim_cube = self.compute_sim_cube(seq1, seq2, truncate=self.classifier_net.input_len)
|
||||
focus_cube = self.compute_focus_cube(sim_cube)
|
||||
logits = self.classifier_net(focus_cube.unsqueeze(0))
|
||||
sim_cube = self.compute_sim_cube(seq1, seq2)
|
||||
truncate = self.classifier_net.input_len
|
||||
if truncate is not None:
|
||||
sim_cube = sim_cube[:, :, :truncate, :truncate].contiguous()
|
||||
pad_cube = pad_cube[:, :truncate, :truncate].contiguous()
|
||||
focus_cube = self.compute_focus_cube(sim_cube, pad_cube)
|
||||
logits = self.classifier_net(focus_cube)
|
||||
return logits
|
||||
|
||||
Executable
+2
@@ -0,0 +1,2 @@
|
||||
#!/bin/sh
|
||||
python . --clip_norm 50 --decay 0.95 --lr 1E-4 --mbatch_size 1 --momentum 0 --optimizer rmsprop --weight_decay 0
|
||||
+5
-6
@@ -12,7 +12,7 @@ class RandomParamIterator(object):
|
||||
return param_set
|
||||
|
||||
class Tuner(object):
|
||||
def __init__(self, *iterators, limit=100):
|
||||
def __init__(self, *iterators, limit=500):
|
||||
self.iterators = iterators
|
||||
self.limit = limit
|
||||
|
||||
@@ -25,14 +25,13 @@ class Tuner(object):
|
||||
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],
|
||||
vgg_param_sets = dict(classifer=["vdpwi"], decay=[0.99, 0.95], lr=[5E-4, 1E-4], mbatch_size=[8, 16],
|
||||
optimizer=["adam", "rmsprop"], weight_decay=[0, 1E-5, 5E-4], momentum=[0, 0.15, 0.05])
|
||||
res_param_sets = dict(classifier=["resnet"], clip_norm=[5, 7, 9], 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()
|
||||
Tuner(vgg_iterator).start()
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user