mirror of
https://github.com/wassname/Castor.git
synced 2026-09-09 11:13:20 +08:00
Add HAN and XML_CNN for Doc Classification (#154)
* Add Reuters option in common.dataset * Add Reuters option in common.dataset * Add HAN model * Add XML-CNN * Add HAN * Add Hierarchical tokenization for Reuters * Add README for HAN * Add XML Readme * Update HAN Readme
This commit is contained in:
committed by
Victor Yang
parent
ed4f01852e
commit
650882fb6e
@@ -69,6 +69,9 @@ class DatasetFactory(object):
|
||||
train_loader, dev_loader, test_loader = PIT2015.iters(dataset_root, word_vectors_file, word_vectors_dir, batch_size, device=device, unk_init=UnknownWordVecCache.unk)
|
||||
embedding = nn.Embedding.from_pretrained(PIT2015.TEXT_FIELD.vocab.vectors)
|
||||
return PIT2015, embedding, train_loader, test_loader, dev_loader
|
||||
|
||||
|
||||
|
||||
elif dataset_name == 'snli':
|
||||
dataset_root = os.path.join(castor_dir, os.pardir, 'Castor-data', 'datasets', 'snli_1.0/')
|
||||
train_loader, dev_loader, test_loader = SNLI.iters(dataset_root, word_vectors_file, word_vectors_dir, batch_size, device=device, unk_init=UnknownWordVecCache.unk)
|
||||
|
||||
+13
-4
@@ -2,7 +2,7 @@ import re
|
||||
import os
|
||||
|
||||
import torch
|
||||
from torchtext.data import Field, TabularDataset
|
||||
from torchtext.data import NestedField ,Field, TabularDataset
|
||||
from torchtext.data.iterator import BucketIterator
|
||||
from torchtext.vocab import Vectors
|
||||
|
||||
@@ -15,6 +15,10 @@ def clean_string(string):
|
||||
string = re.sub(r"\s{2,}", " ", string)
|
||||
return string.lower().strip().split()
|
||||
|
||||
def split_sents(string):
|
||||
string = re.sub(r"[!?]"," ", string)
|
||||
return string.strip().split('.')
|
||||
|
||||
|
||||
def clean_string_fl(string):
|
||||
"""
|
||||
@@ -39,8 +43,8 @@ def process_labels(string):
|
||||
class Reuters(TabularDataset):
|
||||
NAME = 'Reuters'
|
||||
NUM_CLASSES = 90
|
||||
|
||||
TEXT_FIELD = Field(batch_first=True, tokenize=clean_string_fl)
|
||||
|
||||
TEXT_FIELD = Field(batch_first=True, tokenize=clean_string)
|
||||
LABEL_FIELD = Field(sequential=False, use_vocab=False, batch_first=True, preprocessing=process_labels)
|
||||
|
||||
@staticmethod
|
||||
@@ -75,4 +79,9 @@ class Reuters(TabularDataset):
|
||||
train, val, test = cls.splits(path)
|
||||
cls.TEXT_FIELD.build_vocab(train, val, test, vectors=vectors)
|
||||
return BucketIterator.splits((train, val, test), batch_size=batch_size, repeat=False, shuffle=shuffle,
|
||||
sort_within_batch=True, device=device)
|
||||
sort_within_batch=True, device=device)
|
||||
|
||||
class Reuters_hierarchical(Reuters):
|
||||
|
||||
In_FIELD = Field(batch_first = True, tokenize = clean_string)
|
||||
TEXT_FIELD = NestedField(In_FIELD, tokenize = split_sents)
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
# Hierarchical Attention Networks
|
||||
|
||||
Implementation for Hierarchical Attention Networks for Documnet Classification of [HAN (2016)](https://www.cs.cmu.edu/~hovy/papers/16HLT-hierarchical-attention-networks.pdf) with PyTorch and Torchtext.
|
||||
|
||||
## Model Type
|
||||
|
||||
- rand: All words are randomly initialized and then modified during training.
|
||||
- static: A model with pre-trained vectors from [word2vec](https://code.google.com/archive/p/word2vec/). All words -- including the unknown ones that are initialized with zero -- are kept static and only the other parameters of the model are learned.
|
||||
- non-static: Same as above but the pretrained vectors are fine-tuned for each task.
|
||||
|
||||
|
||||
|
||||
## Quick Start
|
||||
|
||||
To run the model on Reuters dataset on static, just run the following from the Castor working directory.
|
||||
|
||||
```
|
||||
python -m han --dataset Reuters
|
||||
```
|
||||
|
||||
The file will be saved in
|
||||
|
||||
```
|
||||
han/saves/best_model.pt
|
||||
```
|
||||
|
||||
To test the model, you can use the following command.
|
||||
|
||||
```
|
||||
python -m han --trained_model han/saves/Reuters/static_best_model.pt
|
||||
```
|
||||
|
||||
## Dataset
|
||||
|
||||
We experiment the model on the following datasets.
|
||||
|
||||
- Reuters-21578: Split the data into sentences for the sentence level attention model and split the sentences into words for the word level attention. The word2vec pretrained embeddings were used for the task.
|
||||
|
||||
## Settings
|
||||
|
||||
Adam is used for training.
|
||||
|
||||
## Training Time
|
||||
|
||||
For training time, when
|
||||
|
||||
```
|
||||
torch.backends.cudnn.deterministic = True
|
||||
```
|
||||
|
||||
is specified, the training will be ~10 min. Reuters-21578 is a relatively small dataset and the implementation is a vectorized one, hence the speed.
|
||||
|
||||
|
||||
|
||||
## TODO
|
||||
- a combined hyperparameter tuning on a few of the datasets and report results with the hyperparameters
|
||||
+189
@@ -0,0 +1,189 @@
|
||||
from copy import deepcopy
|
||||
import logging
|
||||
import random
|
||||
from sklearn import metrics
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.onnx
|
||||
|
||||
from common.evaluation import EvaluatorFactory
|
||||
from common.train import TrainerFactory
|
||||
from datasets.sst import SST1
|
||||
from datasets.sst import SST2
|
||||
from datasets.reuters import Reuters_hierarchical as Reuters
|
||||
from han.args import get_args
|
||||
from han.model import HAN
|
||||
import torch.nn.functional as F
|
||||
|
||||
class UnknownWordVecCache(object):
|
||||
"""
|
||||
Caches the first randomly generated word vector for a certain size to make it is reused.
|
||||
"""
|
||||
cache = {}
|
||||
|
||||
@classmethod
|
||||
def unk(cls, tensor):
|
||||
size_tup = tuple(tensor.size())
|
||||
if size_tup not in cls.cache:
|
||||
cls.cache[size_tup] = torch.Tensor(tensor.size())
|
||||
# choose 0.25 so unknown vectors have approximately same variance as pre-trained ones
|
||||
# same as original implementation: https://github.com/yoonkim/CNN_sentence/blob/0a626a048757d5272a7e8ccede256a434a6529be/process_data.py#L95
|
||||
cls.cache[size_tup].uniform_(-0.25, 0.25)
|
||||
return cls.cache[size_tup]
|
||||
|
||||
|
||||
def get_logger():
|
||||
logger = logging.getLogger(__name__)
|
||||
logger.setLevel(logging.INFO)
|
||||
|
||||
ch = logging.StreamHandler()
|
||||
ch.setLevel(logging.DEBUG)
|
||||
formatter = logging.Formatter('%(levelname)s - %(message)s')
|
||||
ch.setFormatter(formatter)
|
||||
logger.addHandler(ch)
|
||||
|
||||
return logger
|
||||
|
||||
|
||||
def evaluate_dataset(split_name, dataset_cls, model, embedding, loader, batch_size, device):
|
||||
saved_model_evaluator = EvaluatorFactory.get_evaluator(dataset_cls, model, embedding, loader, batch_size, device)
|
||||
scores, metric_names = saved_model_evaluator.get_scores()
|
||||
logger.info('Evaluation metrics for {}'.format(split_name))
|
||||
logger.info('\t'.join([' '] + metric_names))
|
||||
logger.info('\t'.join([split_name] + list(map(str, scores))))
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
# 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:
|
||||
args.gpu = -1
|
||||
if torch.cuda.is_available() and args.cuda:
|
||||
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.')
|
||||
np.random.seed(args.seed)
|
||||
random.seed(args.seed)
|
||||
logger = get_logger()
|
||||
|
||||
# Set up the data for training SST-1
|
||||
if args.dataset == 'SST-1':
|
||||
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, unk_init=UnknownWordVecCache.unk)
|
||||
# Set up the data for training SST-2
|
||||
elif args.dataset == 'SST-2':
|
||||
train_iter, dev_iter, test_iter = SST2.iters(args.data_dir, args.word_vectors_file, args.word_vectors_dir, batch_size=args.batch_size, device=args.gpu, unk_init=UnknownWordVecCache.unk)
|
||||
elif args.dataset == 'Reuters':
|
||||
train_iter, dev_iter, test_iter = Reuters.iters(args.data_dir, args.word_vectors_file, args.word_vectors_dir, batch_size=args.batch_size, device=args.gpu, unk_init=UnknownWordVecCache.unk)
|
||||
else:
|
||||
raise ValueError('Unrecognized dataset')
|
||||
|
||||
config = deepcopy(args)
|
||||
config.dataset = train_iter.dataset
|
||||
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))
|
||||
|
||||
if args.resume_snapshot:
|
||||
if args.cuda:
|
||||
model = torch.load(args.resume_snapshot, map_location=lambda storage, location: storage.cuda(args.gpu))
|
||||
else:
|
||||
model = torch.load(args.resume_snapshot, map_location=lambda storage, location: storage)
|
||||
else:
|
||||
model = HAN(config)
|
||||
if args.cuda:
|
||||
model.cuda()
|
||||
print('Shift model to GPU')
|
||||
|
||||
parameter = filter(lambda p: p.requires_grad, model.parameters())
|
||||
print(parameter)
|
||||
#optimizer = torch.optim.Adadelta(parameter, lr=args.lr, weight_decay=args.weight_decay)
|
||||
#optimizer = torch.optim.SGD(parameter, lr = args.lr, momentum = 0.9)
|
||||
optimizer = torch.optim.Adam(parameter, lr = args.lr)
|
||||
if args.dataset == 'SST-1':
|
||||
train_evaluator = EvaluatorFactory.get_evaluator(SST1, model, None, train_iter, args.batch_size, args.gpu)
|
||||
test_evaluator = EvaluatorFactory.get_evaluator(SST1, model, None, test_iter, args.batch_size, args.gpu)
|
||||
dev_evaluator = EvaluatorFactory.get_evaluator(SST1, model, None, dev_iter, args.batch_size, args.gpu)
|
||||
elif args.dataset == 'SST-2':
|
||||
train_evaluator = EvaluatorFactory.get_evaluator(SST2, model, None, train_iter, args.batch_size, args.gpu)
|
||||
test_evaluator = EvaluatorFactory.get_evaluator(SST2, model, None, test_iter, args.batch_size, args.gpu)
|
||||
dev_evaluator = EvaluatorFactory.get_evaluator(SST2, model, None, dev_iter, args.batch_size, args.gpu)
|
||||
elif args.dataset == 'Reuters':
|
||||
train_evaluator = EvaluatorFactory.get_evaluator(Reuters, model, None, train_iter, args.batch_size, args.gpu)
|
||||
test_evaluator = EvaluatorFactory.get_evaluator(Reuters, model, None, test_iter, args.batch_size, args.gpu)
|
||||
dev_evaluator = EvaluatorFactory.get_evaluator(Reuters, model, None, dev_iter, args.batch_size, args.gpu)
|
||||
else:
|
||||
raise ValueError('Unrecognized dataset')
|
||||
|
||||
trainer_config = {
|
||||
'optimizer': optimizer,
|
||||
'batch_size': args.batch_size,
|
||||
'log_interval': args.log_every,
|
||||
'dev_log_interval': args.dev_every,
|
||||
'patience': args.patience,
|
||||
'model_outfile': args.save_path, # actually a directory, using model_outfile to conform to Trainer naming convention
|
||||
'logger': logger
|
||||
}
|
||||
trainer = TrainerFactory.get_trainer(args.dataset, model, None, train_iter, trainer_config, train_evaluator, test_evaluator, dev_evaluator)
|
||||
|
||||
if not args.trained_model:
|
||||
trainer.train(args.epochs)
|
||||
else:
|
||||
if args.cuda:
|
||||
model = torch.load(args.trained_model, map_location=lambda storage, location: storage.cuda(args.gpu))
|
||||
else:
|
||||
model = torch.load(args.trained_model, map_location=lambda storage, location: storage)
|
||||
|
||||
if args.dataset == 'SST-1':
|
||||
evaluate_dataset('dev', SST1, model, None, dev_iter, args.batch_size, args.gpu)
|
||||
evaluate_dataset('test', SST1, model, None, test_iter, args.batch_size, args.gpu)
|
||||
elif args.dataset == 'SST-2':
|
||||
evaluate_dataset('dev', SST2, model, None, dev_iter, args.batch_size, args.gpu)
|
||||
evaluate_dataset('test', SST2, model, None, test_iter, args.batch_size, args.gpu)
|
||||
elif args.dataset == 'Reuters':
|
||||
evaluate_dataset('dev', Reuters, model, None, dev_iter, args.batch_size, args.gpu)
|
||||
evaluate_dataset('test', Reuters, model, None, test_iter, args.batch_size, args.gpu)
|
||||
else:
|
||||
raise ValueError('Unrecognized dataset')
|
||||
|
||||
|
||||
|
||||
# Calculate dev and test metrics
|
||||
for data_loader in [dev_iter, test_iter]:
|
||||
predicted_labels = list()
|
||||
target_labels = list()
|
||||
for batch_idx, batch in enumerate(data_loader):
|
||||
scores_rounded = F.sigmoid(model(batch.text)).round().long()
|
||||
predicted_labels.extend(scores_rounded.cpu().detach().numpy())
|
||||
target_labels.extend(batch.label.cpu().detach().numpy())
|
||||
predicted_labels = np.array(predicted_labels)
|
||||
target_labels = np.array(target_labels)
|
||||
accuracy = metrics.accuracy_score(target_labels, predicted_labels)
|
||||
precision = metrics.precision_score(target_labels, predicted_labels, average='micro')
|
||||
recall = metrics.recall_score(target_labels, predicted_labels, average='micro')
|
||||
f1 = metrics.f1_score(target_labels, predicted_labels, average='micro')
|
||||
if data_loader == dev_iter:
|
||||
print("Dev metrics:")
|
||||
else:
|
||||
print("Test metrics:")
|
||||
print(accuracy, precision, recall, f1)
|
||||
|
||||
|
||||
|
||||
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 = 'han_{}.onnx'.format(args.mode)
|
||||
torch.onnx.export(model, dummy_input, onnx_filename)
|
||||
print('Exported model in ONNX format as {}'.format(onnx_filename))
|
||||
+44
@@ -0,0 +1,44 @@
|
||||
import os
|
||||
|
||||
from argparse import ArgumentParser
|
||||
|
||||
|
||||
def get_args():
|
||||
parser = ArgumentParser(description="HAN")
|
||||
parser.add_argument('--no_cuda', action='store_false', help='do not use cuda', dest='cuda')
|
||||
parser.add_argument('--gpu', type=int, default=0) # Use -1 for CPU
|
||||
parser.add_argument('--epochs', type=int, default=30)
|
||||
|
||||
|
||||
parser.add_argument('--word_num_hidden', type = int, default = 50)
|
||||
parser.add_argument('--sentence_num_hidden', type = int, default = 50)
|
||||
|
||||
|
||||
parser.add_argument('--batch_size', type=int, default=64)
|
||||
parser.add_argument('--mode', type=str, default='static', choices=['rand', 'static', 'non-static'])
|
||||
parser.add_argument('--lr', type=float, default=1.0)
|
||||
parser.add_argument('--seed', type=int, default=3435)
|
||||
parser.add_argument('--dataset', type=str, default='SST-1', choices=['SST-1', 'SST-2', 'Reuters'])
|
||||
parser.add_argument('--resume_snapshot', type=str, default=None)
|
||||
parser.add_argument('--dev_every', type=int, default=30)
|
||||
parser.add_argument('--log_every', type=int, default=10)
|
||||
parser.add_argument('--patience', type=int, default=50)
|
||||
parser.add_argument('--save_path', type=str, default='han/saves')
|
||||
parser.add_argument('--output_channel', type=int, default=100)
|
||||
parser.add_argument('--words_dim', type=int, default=300)
|
||||
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('--data_dir', help='word vectors directory',
|
||||
default=os.path.join(os.pardir, 'Castor-data', 'datasets'))
|
||||
parser.add_argument('--word_vectors_dir', help='word vectors directory',
|
||||
default=os.path.join(os.pardir, 'Castor-data', 'embeddings', 'word2vec'))
|
||||
parser.add_argument('--word_vectors_file', help='word vectors filename', default='/data/GoogleNews-vectors-negative300.txt')
|
||||
parser.add_argument('--trained_model', type=str, default="")
|
||||
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
|
||||
Executable
+28
@@ -0,0 +1,28 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.autograd import Variable
|
||||
#from utils import
|
||||
import torch.nn.functional as F
|
||||
from han.sent_level_rnn import SentLevelRNN
|
||||
from han.word_level_rnn import WordLevelRNN
|
||||
|
||||
|
||||
class HAN(nn.Module):
|
||||
def __init__(self, config):
|
||||
super(HAN, self).__init__()
|
||||
self.dataset = config.dataset
|
||||
self.mode = config.mode
|
||||
self.word_attention_rnn = WordLevelRNN(config)
|
||||
self.sentence_attention_rnn = SentLevelRNN(config)
|
||||
def forward(self,x):
|
||||
x = x.permute(1,2,0) ## Expected : #sentences, #words, batch size
|
||||
num_sentences = x.size()[0]
|
||||
word_attentions = None
|
||||
for i in range(num_sentences):
|
||||
_word_attention = self.word_attention_rnn(x[i,:,:])
|
||||
if word_attentions is None:
|
||||
word_attentions = _word_attention
|
||||
else:
|
||||
word_attentions = torch.cat((word_attentions, _word_attention),0)
|
||||
return self.sentence_attention_rnn(word_attentions)
|
||||
|
||||
@@ -0,0 +1,32 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.autograd import Variable
|
||||
|
||||
import torch.nn.functional as F
|
||||
|
||||
class SentLevelRNN(nn.Module):
|
||||
def __init__(self, config):
|
||||
super(SentLevelRNN, self).__init__()
|
||||
dataset = config.dataset
|
||||
sentence_num_hidden = config.sentence_num_hidden
|
||||
word_num_hidden = config.word_num_hidden
|
||||
target_class = config.target_class
|
||||
self.sentence_context_wghts = nn.Parameter(torch.rand(2*sentence_num_hidden, 1))
|
||||
self.sentence_context_wghts.data.uniform_(-0.1, 0.1)
|
||||
self.sentence_GRU = nn.GRU(2*word_num_hidden, sentence_num_hidden, bidirectional = True)
|
||||
self.sentence_linear = nn.Linear(2*sentence_num_hidden, 2*sentence_num_hidden, bias = True)
|
||||
self.fc = nn.Linear(2*sentence_num_hidden , target_class)
|
||||
self.soft_sent = nn.Softmax()
|
||||
self.final_log_soft = F.log_softmax
|
||||
|
||||
def forward(self,x):
|
||||
sentence_h,_ = self.sentence_GRU(x)
|
||||
x = torch.tanh(self.sentence_linear(sentence_h))
|
||||
x = torch.matmul(x, self.sentence_context_wghts)
|
||||
x = x.squeeze()
|
||||
x = self.soft_sent(x.transpose(1,0))
|
||||
x = torch.mul(sentence_h.permute(2,0,1), x.transpose(1,0))
|
||||
x = torch.sum(x,dim = 1).transpose(1,0).unsqueeze(0)
|
||||
#x = self.final_log_soft(self.fc(x.squeeze(0)))
|
||||
x = self.fc(x.squeeze(0))
|
||||
return x
|
||||
@@ -0,0 +1,49 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.autograd import Variable
|
||||
import torch.nn.functional as F
|
||||
|
||||
class WordLevelRNN(nn.Module):
|
||||
def __init__(self, config):
|
||||
super(WordLevelRNN, self).__init__()
|
||||
dataset = config.dataset
|
||||
word_num_hidden = config.word_num_hidden
|
||||
words_num = config.words_num
|
||||
words_dim = config.words_dim
|
||||
self.mode = config.mode
|
||||
if self.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 self.mode == 'static':
|
||||
self.static_embed = nn.Embedding.from_pretrained(dataset.TEXT_FIELD.vocab.vectors, freeze = True)
|
||||
elif self.mode == 'non-static':
|
||||
self.non_static_embed = nn.Embedding.from_pretrained(dataset.TEXT_FIELD.vocab.vectors, freeze = False)
|
||||
else:
|
||||
print("Unsupported order")
|
||||
exit()
|
||||
self.word_context_wghts = nn.Parameter(torch.rand(2*word_num_hidden,1))
|
||||
self.GRU = nn.GRU(words_dim, word_num_hidden, bidirectional = True)
|
||||
self.linear = nn.Linear(2*word_num_hidden, 2*word_num_hidden, bias = True)
|
||||
self.word_context_wghts.data.uniform_(-0.25, 0.25)
|
||||
self.soft_word = nn.Softmax()
|
||||
def forward(self, x):
|
||||
##################
|
||||
## x expected to be of dimensions--> (num_words, batch_size)
|
||||
if self.mode == 'rand':
|
||||
x = self.embed(x)
|
||||
elif self.mode == 'static':
|
||||
x = self.static_embed(x)
|
||||
elif self.mode == 'non-static':
|
||||
x = self.non_static_embed(x)
|
||||
else :
|
||||
print("Unsuported mode")
|
||||
exit()
|
||||
h,_ = self.GRU(x)
|
||||
x = torch.tanh(self.linear(h))
|
||||
x = torch.matmul(x, self.word_context_wghts)
|
||||
x = x.squeeze()
|
||||
x = self.soft_word(x.transpose(1,0))
|
||||
x = torch.mul(h.permute(2,0,1), x.transpose(1,0))
|
||||
x = torch.sum(x, dim = 1).transpose(1,0).unsqueeze(0)
|
||||
|
||||
return x
|
||||
@@ -0,0 +1,40 @@
|
||||
# XML_CNN
|
||||
|
||||
Implementation for XML Convolutional Neural Network for Document Classification of [XML-CNN (2014)](http://nyc.lti.cs.cmu.edu/yiming/Publications/jliu-sigir17.pdf) with PyTorch and Torchtext.
|
||||
|
||||
## Model Type
|
||||
|
||||
- rand: All words are randomly initialized and then modified during training.
|
||||
- static: A model with pre-trained vectors from [word2vec](https://code.google.com/archive/p/word2vec/). All words -- including the unknown ones that are initialized with zero -- are kept static and only the other parameters of the model are learned.
|
||||
- non-static: Same as above but the pretrained vectors are fine-tuned for each task.
|
||||
|
||||
## Quick Start
|
||||
|
||||
To run the model on Reuters dataset on static just run the following from the Castor working directory.
|
||||
|
||||
```
|
||||
python -m xml_cnn --dataset Reuters
|
||||
```
|
||||
|
||||
The file will be saved in
|
||||
|
||||
```
|
||||
xml_cnn/saves/best_model.pt
|
||||
```
|
||||
|
||||
|
||||
|
||||
## Dataset
|
||||
|
||||
We experiment the model on the following datasets.
|
||||
|
||||
- Reuters: A multi-label document classification dataset.
|
||||
|
||||
## Settings
|
||||
|
||||
Adam is used for training.
|
||||
|
||||
|
||||
## TODO
|
||||
|
||||
- Report hyperparameters and results after finetuning on other datasets like AAPD.
|
||||
@@ -0,0 +1,186 @@
|
||||
from copy import deepcopy
|
||||
import logging
|
||||
import random
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.onnx
|
||||
import torch.nn.functional as F
|
||||
from sklearn import metrics
|
||||
|
||||
from common.evaluation import EvaluatorFactory
|
||||
from common.train import TrainerFactory
|
||||
from datasets.sst import SST1
|
||||
from datasets.sst import SST2
|
||||
from datasets.reuters import Reuters
|
||||
from xml_cnn.args import get_args
|
||||
from xml_cnn.model import XmlCNN
|
||||
|
||||
|
||||
class UnknownWordVecCache(object):
|
||||
"""
|
||||
Caches the first randomly generated word vector for a certain size to make it is reused.
|
||||
"""
|
||||
cache = {}
|
||||
|
||||
@classmethod
|
||||
def unk(cls, tensor):
|
||||
size_tup = tuple(tensor.size())
|
||||
if size_tup not in cls.cache:
|
||||
cls.cache[size_tup] = torch.Tensor(tensor.size())
|
||||
# choose 0.25 so unknown vectors have approximately same variance as pre-trained ones
|
||||
# same as original implementation: https://github.com/yoonkim/CNN_sentence/blob/0a626a048757d5272a7e8ccede256a434a6529be/process_data.py#L95
|
||||
cls.cache[size_tup].uniform_(-0.25, 0.25)
|
||||
return cls.cache[size_tup]
|
||||
|
||||
|
||||
def get_logger():
|
||||
logger = logging.getLogger(__name__)
|
||||
logger.setLevel(logging.INFO)
|
||||
|
||||
ch = logging.StreamHandler()
|
||||
ch.setLevel(logging.DEBUG)
|
||||
formatter = logging.Formatter('%(levelname)s - %(message)s')
|
||||
ch.setFormatter(formatter)
|
||||
logger.addHandler(ch)
|
||||
|
||||
return logger
|
||||
|
||||
|
||||
def evaluate_dataset(split_name, dataset_cls, model, embedding, loader, batch_size, device):
|
||||
saved_model_evaluator = EvaluatorFactory.get_evaluator(dataset_cls, model, embedding, loader, batch_size, device)
|
||||
scores, metric_names = saved_model_evaluator.get_scores()
|
||||
logger.info('Evaluation metrics for {}'.format(split_name))
|
||||
logger.info('\t'.join([' '] + metric_names))
|
||||
logger.info('\t'.join([split_name] + list(map(str, scores))))
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
# 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:
|
||||
args.gpu = -1
|
||||
if torch.cuda.is_available() and args.cuda:
|
||||
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.')
|
||||
np.random.seed(args.seed)
|
||||
random.seed(args.seed)
|
||||
logger = get_logger()
|
||||
|
||||
# Set up the data for training SST-1
|
||||
if args.dataset == 'SST-1':
|
||||
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, unk_init=UnknownWordVecCache.unk)
|
||||
# Set up the data for training SST-2
|
||||
elif args.dataset == 'SST-2':
|
||||
train_iter, dev_iter, test_iter = SST2.iters(args.data_dir, args.word_vectors_file, args.word_vectors_dir, batch_size=args.batch_size, device=args.gpu, unk_init=UnknownWordVecCache.unk)
|
||||
elif args.dataset == 'Reuters':
|
||||
train_iter, dev_iter, test_iter = Reuters.iters(args.data_dir, args.word_vectors_file, args.word_vectors_dir, batch_size=args.batch_size, device=args.gpu, unk_init=UnknownWordVecCache.unk)
|
||||
else:
|
||||
raise ValueError('Unrecognized dataset')
|
||||
|
||||
config = deepcopy(args)
|
||||
config.dataset = train_iter.dataset
|
||||
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))
|
||||
|
||||
if args.resume_snapshot:
|
||||
if args.cuda:
|
||||
model = torch.load(args.resume_snapshot, map_location=lambda storage, location: storage.cuda(args.gpu))
|
||||
else:
|
||||
model = torch.load(args.resume_snapshot, map_location=lambda storage, location: storage)
|
||||
else:
|
||||
model = XmlCNN(config)
|
||||
if args.cuda:
|
||||
model.cuda()
|
||||
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)
|
||||
optimizer = torch.optim.Adam(parameter, lr = args.lr)
|
||||
|
||||
if args.dataset == 'SST-1':
|
||||
train_evaluator = EvaluatorFactory.get_evaluator(SST1, model, None, train_iter, args.batch_size, args.gpu)
|
||||
test_evaluator = EvaluatorFactory.get_evaluator(SST1, model, None, test_iter, args.batch_size, args.gpu)
|
||||
dev_evaluator = EvaluatorFactory.get_evaluator(SST1, model, None, dev_iter, args.batch_size, args.gpu)
|
||||
elif args.dataset == 'SST-2':
|
||||
train_evaluator = EvaluatorFactory.get_evaluator(SST2, model, None, train_iter, args.batch_size, args.gpu)
|
||||
test_evaluator = EvaluatorFactory.get_evaluator(SST2, model, None, test_iter, args.batch_size, args.gpu)
|
||||
dev_evaluator = EvaluatorFactory.get_evaluator(SST2, model, None, dev_iter, args.batch_size, args.gpu)
|
||||
elif args.dataset == 'Reuters':
|
||||
train_evaluator = EvaluatorFactory.get_evaluator(Reuters, model, None, train_iter, args.batch_size, args.gpu)
|
||||
test_evaluator = EvaluatorFactory.get_evaluator(Reuters, model, None, test_iter, args.batch_size, args.gpu)
|
||||
dev_evaluator = EvaluatorFactory.get_evaluator(Reuters, model, None, dev_iter, args.batch_size, args.gpu)
|
||||
else:
|
||||
raise ValueError('Unrecognized dataset')
|
||||
|
||||
trainer_config = {
|
||||
'optimizer': optimizer,
|
||||
'batch_size': args.batch_size,
|
||||
'log_interval': args.log_every,
|
||||
'dev_log_interval': args.dev_every,
|
||||
'patience': args.patience,
|
||||
'model_outfile': args.save_path, # actually a directory, using model_outfile to conform to Trainer naming convention
|
||||
'logger': logger
|
||||
}
|
||||
trainer = TrainerFactory.get_trainer(args.dataset, model, None, train_iter, trainer_config, train_evaluator, test_evaluator, dev_evaluator)
|
||||
|
||||
if not args.trained_model:
|
||||
trainer.train(args.epochs)
|
||||
else:
|
||||
if args.cuda:
|
||||
model = torch.load(args.trained_model, map_location=lambda storage, location: storage.cuda(args.gpu))
|
||||
else:
|
||||
model = torch.load(args.trained_model, map_location=lambda storage, location: storage)
|
||||
|
||||
if args.dataset == 'SST-1':
|
||||
evaluate_dataset('dev', SST1, model, None, dev_iter, args.batch_size, args.gpu)
|
||||
evaluate_dataset('test', SST1, model, None, test_iter, args.batch_size, args.gpu)
|
||||
elif args.dataset == 'SST-2':
|
||||
evaluate_dataset('dev', SST2, model, None, dev_iter, args.batch_size, args.gpu)
|
||||
evaluate_dataset('test', SST2, model, None, test_iter, args.batch_size, args.gpu)
|
||||
elif args.dataset == 'Reuters':
|
||||
evaluate_dataset('dev', Reuters, model, None, dev_iter, args.batch_size, args.gpu)
|
||||
evaluate_dataset('test', Reuters, model, None, test_iter, args.batch_size, args.gpu)
|
||||
else:
|
||||
raise ValueError('Unrecognized dataset')
|
||||
|
||||
# Calculate dev and test metrics
|
||||
for data_loader in [dev_iter, test_iter]:
|
||||
predicted_labels = list()
|
||||
target_labels = list()
|
||||
for batch_idx, batch in enumerate(data_loader):
|
||||
scores_rounded = F.sigmoid(model(batch.text)).round().long()
|
||||
predicted_labels.extend(scores_rounded.cpu().detach().numpy())
|
||||
target_labels.extend(batch.label.cpu().detach().numpy())
|
||||
predicted_labels = np.array(predicted_labels)
|
||||
target_labels = np.array(target_labels)
|
||||
accuracy = metrics.accuracy_score(target_labels, predicted_labels)
|
||||
precision = metrics.precision_score(target_labels, predicted_labels, average='micro')
|
||||
recall = metrics.recall_score(target_labels, predicted_labels, average='micro')
|
||||
f1 = metrics.f1_score(target_labels, predicted_labels, average='micro')
|
||||
if data_loader == dev_iter:
|
||||
print("Dev metrics:")
|
||||
else:
|
||||
print("Test metrics:")
|
||||
print(accuracy, precision, recall, f1)
|
||||
|
||||
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 = 'xmlcnn_{}.onnx'.format(args.mode)
|
||||
torch.onnx.export(model, dummy_input, onnx_filename)
|
||||
print('Exported model in ONNX format as {}'.format(onnx_filename))
|
||||
@@ -0,0 +1,43 @@
|
||||
import os
|
||||
|
||||
from argparse import ArgumentParser
|
||||
|
||||
|
||||
def get_args():
|
||||
parser = ArgumentParser(description="Kim CNN")
|
||||
parser.add_argument('--no_cuda', action='store_false', help='do not use cuda', dest='cuda')
|
||||
parser.add_argument('--gpu', type=int, default=0) # Use -1 for CPU
|
||||
parser.add_argument('--epochs', type=int, default=30)
|
||||
parser.add_argument('--batch_size', type=int, default=1024)
|
||||
parser.add_argument('--mode', type=str, default='multichannel', choices=['rand', 'static', 'non-static', 'multichannel'])
|
||||
parser.add_argument('--lr', type=float, default=1.0)
|
||||
parser.add_argument('--seed', type=int, default=3435)
|
||||
parser.add_argument('--dataset', type=str, default='SST-1', choices=['SST-1', 'SST-2', 'Reuters'])
|
||||
parser.add_argument('--resume_snapshot', type=str, default=None)
|
||||
parser.add_argument('--dev_every', type=int, default=30)
|
||||
parser.add_argument('--log_every', type=int, default=10)
|
||||
parser.add_argument('--patience', type=int, default=50)
|
||||
parser.add_argument('--save_path', type=str, default='kim_cnn/saves')
|
||||
parser.add_argument('--output_channel', type=int, default=100)
|
||||
parser.add_argument('--words_dim', type=int, default=300)
|
||||
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('--num_bottleneck_hidden', type=int, default=512) #bottleneck layer
|
||||
parser.add_argument('--dynamic_pool_length', type=int, default=32) #dynamic pool length
|
||||
|
||||
|
||||
parser.add_argument('--data_dir', help='word vectors directory',
|
||||
default=os.path.join(os.pardir, 'Castor-data', 'datasets'))
|
||||
parser.add_argument('--word_vectors_dir', help='word vectors directory',
|
||||
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('--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
|
||||
@@ -0,0 +1,76 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class XmlCNN(nn.Module):
|
||||
def __init__(self, config):
|
||||
super(XmlCNN, self).__init__()
|
||||
dataset = config.dataset
|
||||
self.output_channel = config.output_channel
|
||||
target_class = config.target_class
|
||||
words_num = config.words_num
|
||||
words_dim = config.words_dim
|
||||
self.mode = config.mode
|
||||
self.num_bottleneck_hidden = config.num_bottleneck_hidden
|
||||
self.dynamic_pool_length = config.dynamic_pool_length
|
||||
self.Ks = 3 # There are three conv nets here
|
||||
|
||||
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:
|
||||
print("Unsupported Mode")
|
||||
exit()
|
||||
|
||||
## Different filter sizes in xml_cnn than kim_cnn
|
||||
|
||||
self.conv1 = nn.Conv2d(input_channel, self.output_channel, (2, words_dim), padding=(1,0))
|
||||
self.conv2 = nn.Conv2d(input_channel, self.output_channel, (4, words_dim), padding=(3,0))
|
||||
self.conv3 = nn.Conv2d(input_channel, self.output_channel, (8, words_dim), padding=(7,0))
|
||||
|
||||
|
||||
self.dropout = nn.Dropout(config.dropout)
|
||||
self.bottleneck = nn.Linear(self.Ks*self.output_channel*self.dynamic_pool_length, self.num_bottleneck_hidden)
|
||||
self.fc1 = nn.Linear(self.num_bottleneck_hidden, target_class)
|
||||
|
||||
self.pool = nn.AdaptiveMaxPool1d(self.dynamic_pool_length) #Adaptive pooling
|
||||
|
||||
|
||||
|
||||
def forward(self, x):
|
||||
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)
|
||||
elif self.mode == 'static':
|
||||
static_input = self.static_embed(x)
|
||||
x = static_input.unsqueeze(1) # (batch, channel_input, sent_len, embed_dim)
|
||||
elif self.mode == 'non-static':
|
||||
non_static_input = self.non_static_embed(x)
|
||||
x = non_static_input.unsqueeze(1) # (batch, channel_input, sent_len, embed_dim)
|
||||
elif self.mode == 'multichannel':
|
||||
non_static_input = self.non_static_embed(x)
|
||||
static_input = self.static_embed(x)
|
||||
x = torch.stack([non_static_input, static_input], dim=1) # (batch, channel_input=2, sent_len, embed_dim)
|
||||
else:
|
||||
print("Unsupported Mode")
|
||||
exit()
|
||||
x = [F.relu(self.conv1(x)).squeeze(3), F.relu(self.conv2(x)).squeeze(3), F.relu(self.conv3(x)).squeeze(3)]
|
||||
x = [self.pool(i).squeeze(2) for i in x]
|
||||
|
||||
# (batch, channel_output) * Ks
|
||||
x = torch.cat(x, 1) # (batch, channel_output * Ks)
|
||||
x = F.relu(self.bottleneck(x.view(-1, self.Ks*self.output_channel*self.dynamic_pool_length)))
|
||||
x = self.dropout(x)
|
||||
logit = self.fc1(x) # (batch, target_size)
|
||||
return logit
|
||||
Reference in New Issue
Block a user