mirror of
https://github.com/wassname/Castor.git
synced 2026-09-10 11:40:44 +08:00
Replication of STOA for Reuters Dataset (#152)
* Add ReutersTrainer, ReutersEvaluator options in Factory classes * Add Reuters to Kim-CNN command line arguments * Fix SST dataset path according to changes in Kim-CNN args The dataset path in args.py was made to point at the dataset folder rather than dataset/SST folder. Hence SST folder was added to paths in the SST dataset class * Add Reuters dataset class, and support in __main__ * Add Reuters dataset trainers and evaluators * Remove debug print statement in reuters_evaluator * Fix rounding bug in reuters_trainer and reuters_evaluator * Add LSTM for baseline text classification measurements * Add eval metrics for lstm_baseline * Set batch_first param in lstm_baseline * Remove onnx args from lstm_baseline * Pack padded sequences in LSTM_baseline * Add TensorBoardX support for Reuters trainer * Add Arxiv Academic Paper Dataset (AAPD) * Add Hidden Bottleneck Layer to BiLSTM * Fix packing of padded tensors in Reuters * Add cmdline args for Hidden Bottleneck Layer for BiLSTM * Include pre-padding lengths in AAPD dataset * Remove duplication of preprocessing code in AAPD * Remove batch_size condition in ReutersTrainer
This commit is contained in:
@@ -0,0 +1,59 @@
|
||||
import re
|
||||
import os
|
||||
|
||||
import torch
|
||||
from datasets.reuters import clean_string, clean_string_fl
|
||||
from torchtext.data import Field, TabularDataset
|
||||
from torchtext.data.iterator import BucketIterator
|
||||
from torchtext.vocab import Vectors
|
||||
|
||||
|
||||
def process_labels(string):
|
||||
"""
|
||||
Returns the label string as a list of integers
|
||||
:param string:
|
||||
:return:
|
||||
"""
|
||||
return [float(x) for x in string]
|
||||
|
||||
|
||||
class AAPD(TabularDataset):
|
||||
NAME = 'AAPD'
|
||||
NUM_CLASSES = 54
|
||||
|
||||
TEXT_FIELD = Field(batch_first=True, tokenize=clean_string, include_lengths=True)
|
||||
LABEL_FIELD = Field(sequential=False, use_vocab=False, batch_first=True, preprocessing=process_labels)
|
||||
|
||||
@staticmethod
|
||||
def sort_key(ex):
|
||||
return len(ex.text)
|
||||
|
||||
@classmethod
|
||||
def splits(cls, path, train=os.path.join('AAPD', 'data', 'aapd_train.tsv'),
|
||||
validation=os.path.join('AAPD', 'data', 'aapd_validation.tsv'),
|
||||
test=os.path.join('AAPD', 'data','aapd_test.tsv'), **kwargs):
|
||||
return super(AAPD, 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, vectors=vectors)
|
||||
return BucketIterator.splits((train, val, test), batch_size=batch_size, repeat=False, shuffle=shuffle,
|
||||
sort_within_batch=True, device=device)
|
||||
Reference in New Issue
Block a user