From 8a00f9cdcdee682a630693aef8cfc1b956bf4f07 Mon Sep 17 00:00:00 2001 From: Michael Tu Date: Sat, 4 Aug 2018 17:30:24 -0400 Subject: [PATCH] Make Kim CNN ONNX-exportable (#136) * Kim CNN - only set embedding for corresponding mode * Kim CNN ONNX Export * Specify dummy ONNX input size from command line --- .gitignore | 1 + common/evaluators/sst_evaluator.py | 2 +- common/trainers/sst_trainer.py | 2 +- kim_cnn/__main__.py | 25 ++++++++++++++++--------- kim_cnn/args.py | 5 ++++- kim_cnn/model.py | 21 ++++++++++++++------- 6 files changed, 37 insertions(+), 19 deletions(-) diff --git a/.gitignore b/.gitignore index e371d91..27220b3 100644 --- a/.gitignore +++ b/.gitignore @@ -12,3 +12,4 @@ text/ kim_cnn/data .results .qrel +*.onnx diff --git a/common/evaluators/sst_evaluator.py b/common/evaluators/sst_evaluator.py index 77cf172..f6368b8 100644 --- a/common/evaluators/sst_evaluator.py +++ b/common/evaluators/sst_evaluator.py @@ -13,7 +13,7 @@ class SSTEvaluator(Evaluator): total_loss = 0 for batch_idx, batch in enumerate(self.data_loader): - scores = self.model(batch) + scores = self.model(batch.text) n_dev_correct += ( torch.max(scores, 1)[1].view(batch.label.size()).data == batch.label.data).sum().item() total_loss += F.cross_entropy(scores, batch.label, size_average=False).item() diff --git a/common/trainers/sst_trainer.py b/common/trainers/sst_trainer.py index 7d4b298..076dda2 100644 --- a/common/trainers/sst_trainer.py +++ b/common/trainers/sst_trainer.py @@ -28,7 +28,7 @@ class SSTTrainer(Trainer): self.iterations += 1 self.model.train() self.optimizer.zero_grad() - scores = self.model(batch) + scores = self.model(batch.text) n_correct += (torch.max(scores, 1)[1].view(batch.label.size()).data == batch.label.data).sum().item() n_total += batch.batch_size train_acc = 100. * n_correct / n_total diff --git a/kim_cnn/__main__.py b/kim_cnn/__main__.py index e998dd0..d1c76f0 100644 --- a/kim_cnn/__main__.py +++ b/kim_cnn/__main__.py @@ -4,6 +4,7 @@ import random import numpy as np import torch +import torch.onnx from common.evaluation import EvaluatorFactory from common.train import TrainerFactory @@ -60,11 +61,11 @@ if __name__ == '__main__': if not args.cuda: args.gpu = -1 if torch.cuda.is_available() and args.cuda: - print("Note: You are using GPU for training") + print('Note: You are using GPU for training') torch.cuda.set_device(args.gpu) torch.cuda.manual_seed(args.seed) if torch.cuda.is_available() and not args.cuda: - print("Warning: You have Cuda but not use it. You are using CPU for training.") + print('Warning: You have Cuda but not use it. You are using CPU for training.') np.random.seed(args.seed) random.seed(args.seed) logger = get_logger() @@ -83,12 +84,12 @@ if __name__ == '__main__': config.target_class = train_iter.dataset.NUM_CLASSES config.words_num = len(train_iter.dataset.TEXT_FIELD.vocab) - print("Dataset {} Mode {}".format(args.dataset, args.mode)) - print("VOCAB num",len(train_iter.dataset.TEXT_FIELD.vocab)) - print("LABEL.target_class:", train_iter.dataset.NUM_CLASSES) - print("Train instance", len(train_iter.dataset)) - print("Dev instance", len(dev_iter.dataset)) - print("Test instance", len(test_iter.dataset)) + print('Dataset {} Mode {}'.format(args.dataset, args.mode)) + print('VOCAB num',len(train_iter.dataset.TEXT_FIELD.vocab)) + print('LABEL.target_class:', train_iter.dataset.NUM_CLASSES) + print('Train instance', len(train_iter.dataset)) + print('Dev instance', len(dev_iter.dataset)) + print('Test instance', len(test_iter.dataset)) if args.resume_snapshot: if args.cuda: @@ -99,7 +100,7 @@ if __name__ == '__main__': model = KimCNN(config) if args.cuda: model.cuda() - print("Shift model to GPU") + print('Shift model to GPU') parameter = filter(lambda p: p.requires_grad, model.parameters()) optimizer = torch.optim.Adadelta(parameter, lr=args.lr, weight_decay=args.weight_decay) @@ -143,3 +144,9 @@ if __name__ == '__main__': else: raise ValueError('Unrecognized dataset') + if args.onnx: + device = torch.device('cuda') if torch.cuda.is_available() and args.cuda else torch.device('cpu') + dummy_input = torch.zeros(args.onnx_batch_size, args.onnx_sent_len, dtype=torch.long, device=device) + onnx_filename = 'kimcnn_{}.onnx'.format(args.mode) + torch.onnx.export(model, dummy_input, onnx_filename) + print('Exported model in ONNX format as {}'.format(onnx_filename)) diff --git a/kim_cnn/args.py b/kim_cnn/args.py index 27fde0f..1c42b6e 100644 --- a/kim_cnn/args.py +++ b/kim_cnn/args.py @@ -29,7 +29,10 @@ def get_args(): default=os.path.join(os.pardir, 'Castor-data', 'embeddings', 'word2vec')) parser.add_argument('--word_vectors_file', help='word vectors filename', default='GoogleNews-vectors-negative300.txt') parser.add_argument('--trained_model', type=str, default="") - parser.add_argument('--weight_decay',type=float, default=0) + parser.add_argument('--weight_decay', type=float, default=0) + parser.add_argument('--onnx', action='store_true', default=False, help='Export model in ONNX format') + parser.add_argument('--onnx_batch_size', type=int, default=1024, help='Batch size for ONNX export') + parser.add_argument('--onnx_sent_len', type=int, default=32, help='Sentence length for ONNX export') args = parser.parse_args() return args diff --git a/kim_cnn/model.py b/kim_cnn/model.py index 384e5b1..1e54520 100644 --- a/kim_cnn/model.py +++ b/kim_cnn/model.py @@ -14,14 +14,22 @@ class KimCNN(nn.Module): words_dim = config.words_dim self.mode = config.mode Ks = 3 # There are three conv nets here - if config.mode == 'multichannel': + + input_channel = 1 + if config.mode == 'rand': + rand_embed_init = torch.Tensor(words_num, words_dim).uniform_(-0.25, 0.25) + self.embed = nn.Embedding.from_pretrained(rand_embed_init, freeze=False) + elif config.mode == 'static': + self.static_embed = nn.Embedding.from_pretrained(dataset.TEXT_FIELD.vocab.vectors, freeze=True) + elif config.mode == 'non-static': + self.non_static_embed = nn.Embedding.from_pretrained(dataset.TEXT_FIELD.vocab.vectors, freeze=False) + elif config.mode == 'multichannel': + self.static_embed = nn.Embedding.from_pretrained(dataset.TEXT_FIELD.vocab.vectors, freeze=True) + self.non_static_embed = nn.Embedding.from_pretrained(dataset.TEXT_FIELD.vocab.vectors, freeze=False) input_channel = 2 else: - input_channel = 1 - rand_embed_init = torch.Tensor(words_num, words_dim).uniform_(-0.25, 0.25) - self.embed = nn.Embedding.from_pretrained(rand_embed_init, freeze=False) - self.static_embed = nn.Embedding.from_pretrained(dataset.TEXT_FIELD.vocab.vectors, freeze=True) - self.non_static_embed = nn.Embedding.from_pretrained(dataset.TEXT_FIELD.vocab.vectors, freeze=False) + print("Unsupported Mode") + exit() self.conv1 = nn.Conv2d(input_channel, output_channel, (3, words_dim), padding=(2,0)) self.conv2 = nn.Conv2d(input_channel, output_channel, (4, words_dim), padding=(3,0)) @@ -31,7 +39,6 @@ class KimCNN(nn.Module): self.fc1 = nn.Linear(Ks * output_channel, target_class) def forward(self, x): - x = x.text if self.mode == 'rand': word_input = self.embed(x) # (batch, sent_len, embed_dim) x = word_input.unsqueeze(1) # (batch, channel_input, sent_len, embed_dim)