diff --git a/datasets/sst.py b/datasets/sst.py new file mode 100644 index 0000000..358372f --- /dev/null +++ b/datasets/sst.py @@ -0,0 +1,57 @@ +import re + +import torch +from torchtext.data import Field, TabularDataset +from torchtext.data.iterator import BucketIterator +from torchtext.vocab import Vectors + + +def clean_str_sst(string): + """ + Tokenization/string cleaning for the SST dataset + """ + string = re.sub(r"[^A-Za-z0-9(),!?\'\`]", " ", string) + string = re.sub(r"\s{2,}", " ", string) + return string.lower().strip().split() + + +class SST1(TabularDataset): + NAME = 'sst-1' + NUM_CLASSES = 5 + + TEXT_FIELD = Field(batch_first=True, tokenize=clean_str_sst) + LABEL_FIELD = Field(sequential=False, use_vocab=False, batch_first=True) + + @staticmethod + def sort_key(ex): + return len(ex.text) + + @classmethod + def splits(cls, path, train='stsa.fine.phrases.train', validation='stsa.fine.dev', test='stsa.fine.test', **kwargs): + return super(SST1, cls).splits( + path, train=train, validation=validation, test=test, + format='tsv', fields=[('label', cls.LABEL_FIELD), ('text', cls.TEXT_FIELD)] + ) + + @classmethod + def iters(cls, path, vectors_name, vectors_cache, batch_size=64, shuffle=True, device=0, vectors=None, + unk_init=torch.Tensor.zero_): + """ + :param path: directory containing train, test, dev files + :param vectors_name: name of word vectors file + :param vectors_cache: path to directory containing word vectors file + :param batch_size: batch size + :param device: GPU device + :param vectors: custom vectors - either predefined torchtext vectors or your own custom Vector classes + :param unk_init: function used to generate vector for OOV words + :return: + """ + if vectors is None: + vectors = Vectors(name=vectors_name, cache=vectors_cache, unk_init=unk_init) + + train, val, test = cls.splits(path) + + cls.TEXT_FIELD.build_vocab(train, val, test, min_freq=2, vectors=vectors) + + return BucketIterator.splits((train, val, test), batch_size=batch_size, repeat=False, shuffle=shuffle, + sort_within_batch=True, device=device) diff --git a/kim_cnn/__init__.py b/kim_cnn/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/kim_cnn/args.py b/kim_cnn/args.py index 6e12269..ebf60a2 100644 --- a/kim_cnn/args.py +++ b/kim_cnn/args.py @@ -22,10 +22,13 @@ def get_args(): parser.add_argument('--embed_dim', type=int, default=300) parser.add_argument('--dropout', type=float, default=0.5) parser.add_argument('--epoch_decay', type=int, default=15) - parser.add_argument('--vector_cache', type=str, default="data/word2vec.sst-1.pt") + parser.add_argument('--data_dir', help='word vectors directory', + default=os.path.join(os.pardir, os.pardir, 'Castor-data', 'datasets', 'SST')) + parser.add_argument('--word_vectors_dir', help='word vectors directory', + default=os.path.join(os.pardir, 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) - args = parser.parse_args() return args diff --git a/kim_cnn/main.py b/kim_cnn/main.py index cd9a7b6..d7f37d9 100644 --- a/kim_cnn/main.py +++ b/kim_cnn/main.py @@ -4,9 +4,8 @@ import numpy as np import torch from torchtext import data from args import get_args -from SST1 import SST1Dataset -from utils import clean_str_sst +from datasets.sst import SST1 args = get_args() torch.manual_seed(args.seed) @@ -26,26 +25,12 @@ if not args.trained_model: sys.exit(1) if args.dataset == 'SST-1': - TEXT = data.Field(batch_first=True, lower=True, tokenize=clean_str_sst) - LABEL = data.Field(sequential=False) - train, dev, test = SST1Dataset.splits(TEXT, LABEL) - -TEXT.build_vocab(train, min_freq=2) -LABEL.build_vocab(train) - -train_iter = data.Iterator(train, batch_size=args.batch_size, device=args.gpu, train=True, repeat=False, - sort=False, shuffle=True) -dev_iter = data.Iterator(dev, batch_size=args.batch_size, device=args.gpu, train=False, repeat=False, - sort=False, shuffle=False) -test_iter = data.Iterator(test, batch_size=args.batch_size, device=args.gpu, train=False, repeat=False, - sort=False, shuffle=False) + train_iter, dev_iter, test_iter = SST1.iters(args.data_dir, args.word_vectors_file, args.word_vectors_dir, batch_size=args.batch_size, device=args.gpu) config = args -config.target_class = len(LABEL.vocab) -config.words_num = len(TEXT.vocab) -config.embed_num = len(TEXT.vocab) - -print("Label dict:", LABEL.vocab.itos) +config.target_class = train_iter.dataset.NUM_CLASSES +config.words_num = len(train_iter.dataset.TEXT_FIELD.vocab) +config.embed_num = len(train_iter.dataset.TEXT_FIELD.vocab) if args.cuda: model = torch.load(args.trained_model, map_location=lambda storage, location: storage.cuda(args.gpu)) @@ -53,7 +38,7 @@ else: model = torch.load(args.trained_model, map_location=lambda storage,location: storage) -def predict(dataset_iter, dataset, dataset_name): +def predict(dataset_iter, dataset_name): print("Dataset: {}".format(dataset_name)) model.eval() dataset_iter.init_epoch() @@ -63,12 +48,12 @@ def predict(dataset_iter, dataset, dataset_name): scores = model(data_batch) n_correct += (torch.max(scores, 1)[1].view(data_batch.label.size()).data == data_batch.label.data).sum() - print("no. correct {} out of {}".format(n_correct, len(dataset))) - accuracy = 100. * n_correct / len(dataset) + print("no. correct {} out of {}".format(n_correct, len(dataset_iter.dataset.examples))) + accuracy = 100. * n_correct / len(dataset_iter.dataset.examples) print("{} accuracy: {:8.6f}%".format(dataset_name, accuracy)) # Run the model on the dev set -predict(dataset_iter=dev_iter, dataset=dev, dataset_name="valid") +predict(dataset_iter=dev_iter, dataset_name="valid") # Run the model on the test set -predict(dataset_iter=test_iter, dataset=test, dataset_name="test") +predict(dataset_iter=test_iter, dataset_name="test") diff --git a/kim_cnn/model.py b/kim_cnn/model.py index d66359d..36f9a9f 100644 --- a/kim_cnn/model.py +++ b/kim_cnn/model.py @@ -3,6 +3,7 @@ import torch.nn as nn import torch.nn.functional as F + class KimCNN(nn.Module): def __init__(self, config): super(KimCNN, self).__init__() @@ -30,7 +31,6 @@ class KimCNN(nn.Module): self.dropout = nn.Dropout(config.dropout) self.fc1 = nn.Linear(Ks * output_channel, target_class) - def forward(self, x): x = x.text if self.mode == 'rand': diff --git a/kim_cnn/train.py b/kim_cnn/train.py index fb8abfc..fddba83 100644 --- a/kim_cnn/train.py +++ b/kim_cnn/train.py @@ -4,17 +4,15 @@ import random import torch import torch.nn as nn import numpy as np -from torchtext import data + +from datasets.sst import SST1 from args import get_args from model import KimCNN -from SST1 import SST1Dataset -from utils import clean_str_sst # Set default configuration in : args.py args = get_args() # Set random seed for reproducibility - torch.manual_seed(args.seed) torch.backends.cudnn.deterministic = True if not args.cuda: @@ -28,54 +26,21 @@ if torch.cuda.is_available() and not args.cuda: np.random.seed(args.seed) random.seed(args.seed) -# Set up the data for training -# SST-1 +# Set up the data for training SST-1 if args.dataset == 'SST-1': - TEXT = data.Field(batch_first=True, tokenize=clean_str_sst) - LABEL = data.Field(sequential=False) - train, dev, test = SST1Dataset.splits(TEXT, LABEL) - -TEXT.build_vocab(train, min_freq=2) -LABEL.build_vocab(train) - -if os.path.isfile(args.vector_cache): - stoi, vectors, dim = torch.load(args.vector_cache) - TEXT.vocab.vectors = torch.Tensor(len(TEXT.vocab), dim) - for i, token in enumerate(TEXT.vocab.itos): - wv_index = stoi.get(token, None) - if wv_index is not None: - TEXT.vocab.vectors[i] = vectors[wv_index] - else: - TEXT.vocab.vectors[i] = torch.Tensor.zero_(TEXT.vocab.vectors[i]) -else: - print("Error: Need word embedding pt file") - exit(1) - -#print('len(TEXT.vocab)', len(TEXT.vocab)) -#print('TEXT.vocab.vectors.size()', TEXT.vocab.vectors.size()) - -train_iter = data.Iterator(train, batch_size=args.batch_size, device=args.gpu, train=True, repeat=False, - sort=False, shuffle=True) -dev_iter = data.Iterator(dev, batch_size=args.batch_size, device=args.gpu, train=False, repeat=False, - sort=False, shuffle=False) -test_iter = data.Iterator(test, batch_size=args.batch_size, device=args.gpu, train=False, repeat=False, - sort=False, shuffle=False) + train_iter, dev_iter, test_iter = SST1.iters(args.data_dir, args.word_vectors_file, args.word_vectors_dir, batch_size=args.batch_size, device=args.gpu) config = args -config.target_class = len(LABEL.vocab) -config.words_num = len(TEXT.vocab) -config.embed_num = len(TEXT.vocab) +config.target_class = train_iter.dataset.NUM_CLASSES +config.words_num = len(train_iter.dataset.TEXT_FIELD.vocab) +config.embed_num = len(train_iter.dataset.TEXT_FIELD.vocab) - -#print(config) print("Dataset {} Mode {}".format(args.dataset, args.mode)) -print("VOCAB num",len(TEXT.vocab)) -print("LABEL.target_class:", len(LABEL.vocab)) -print("LABELS:",LABEL.vocab.itos) -print("Train instance", len(train)) -print("Dev instance", len(dev)) -print("Test instance", len(test)) - +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: @@ -84,8 +49,8 @@ if args.resume_snapshot: model = torch.load(args.resume_snapshot, map_location=lambda storage, location: storage) else: model = KimCNN(config) - model.static_embed.weight.data.copy_(TEXT.vocab.vectors) - model.non_static_embed.weight.data.copy_(TEXT.vocab.vectors) + model.static_embed.weight.data.copy_(train_iter.dataset.TEXT_FIELD.vocab.vectors) + model.non_static_embed.weight.data.copy_(train_iter.dataset.TEXT_FIELD.vocab.vectors) if args.cuda: model.cuda() print("Shift model to GPU") @@ -119,11 +84,9 @@ while True: n_correct, n_total = 0, 0 for batch_idx, batch in enumerate(train_iter): - # Batch size : (Sentence Length, Batch_size) iterations += 1 - model.train(); optimizer.zero_grad() - #print("Text Size:", batch.text.size()) - #print("Label Size:", batch.label.size()) + model.train() + optimizer.zero_grad() scores = model(batch) n_correct += (torch.max(scores, 1)[1].view(batch.label.size()).data == batch.label.data).sum() n_total += batch.batch_size @@ -134,22 +97,22 @@ while True: optimizer.step() - # Evaluate performance on validation set if iterations % args.dev_every == 1: # switch model into evalutaion mode - model.eval(); dev_iter.init_epoch() + model.eval() + dev_iter.init_epoch() n_dev_correct = 0 dev_losses = [] for dev_batch_idx, dev_batch in enumerate(dev_iter): scores = model(dev_batch) n_dev_correct += (torch.max(scores, 1)[1].view(dev_batch.label.size()).data == dev_batch.label.data).sum() dev_loss = criterion(scores, dev_batch.label) - dev_losses.append(dev_loss.data[0]) - dev_acc = 100. * n_dev_correct / len(dev) + dev_losses.append(dev_loss.item()) + dev_acc = 100. * n_dev_correct / len(dev_iter.dataset) print(dev_log_template.format(time.time() - start, epoch, iterations, 1 + batch_idx, len(train_iter), - 100. * (1 + batch_idx) / len(train_iter), loss.data[0], + 100. * (1 + batch_idx) / len(train_iter), loss.item(), sum(dev_losses) / len(dev_losses), train_acc, dev_acc)) # Update validation results @@ -168,25 +131,5 @@ while True: # print progress message print(log_template.format(time.time() - start, epoch, iterations, 1 + batch_idx, len(train_iter), - 100. * (1 + batch_idx) / len(train_iter), loss.data[0], ' ' * 8, + 100. * (1 + batch_idx) / len(train_iter), loss.item(), ' ' * 8, n_correct / n_total * 100, ' ' * 12)) - - - - - - - - - - - - - - - - - - - - diff --git a/kim_cnn/utils.py b/kim_cnn/utils.py index 12ae678..0b21cec 100644 --- a/kim_cnn/utils.py +++ b/kim_cnn/utils.py @@ -19,12 +19,3 @@ def clean_str(string): string = re.sub(r"\?", " ? ", string) string = re.sub(r"\s{2,}", " ", string) return string.lower().strip().split() - - -def clean_str_sst(string): - """ - Tokenization/string cleaning for the SST dataset - """ - string = re.sub(r"[^A-Za-z0-9(),!?\'\`]", " ", string) - string = re.sub(r"\s{2,}", " ", string) - return string.lower().strip().split() \ No newline at end of file