mirror of
https://github.com/wassname/Castor.git
synced 2026-09-09 11:13:20 +08:00
Create Castor-level SST dataset (#123)
* Move SST to root directory datasets * Fix bugs
This commit is contained in:
@@ -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)
|
||||
+5
-2
@@ -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
|
||||
|
||||
+10
-25
@@ -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")
|
||||
|
||||
+1
-1
@@ -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':
|
||||
|
||||
+22
-79
@@ -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))
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user