mirror of
https://github.com/wassname/Castor.git
synced 2026-08-20 12:00:37 +08:00
committed by
Ralph Tang
parent
addc4506d1
commit
97bdaec5bb
+19
-33
@@ -13,6 +13,7 @@ from common.train import TrainerFactory
|
||||
from datasets.sst import SST1
|
||||
from datasets.sst import SST2
|
||||
from datasets.reuters import Reuters
|
||||
from datasets.aapd import AAPD
|
||||
from xml_cnn.args import get_args
|
||||
from xml_cnn.model import XmlCNN
|
||||
|
||||
@@ -74,16 +75,16 @@ if __name__ == '__main__':
|
||||
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:
|
||||
dataset_map = {
|
||||
'SST-1': SST1,
|
||||
'SST-2': SST2,
|
||||
'Reuters': Reuters,
|
||||
'AAPD': AAPD
|
||||
}
|
||||
if args.dataset not in dataset_map:
|
||||
raise ValueError('Unrecognized dataset')
|
||||
else:
|
||||
train_iter, dev_iter, test_iter = dataset_map[args.dataset].iters(args.data_dir, args.word_vectors_file, args.word_vectors_dir, batch_size=args.batch_size, device=args.gpu, unk_init=UnknownWordVecCache.unk)
|
||||
|
||||
config = deepcopy(args)
|
||||
config.dataset = train_iter.dataset
|
||||
@@ -112,21 +113,12 @@ if __name__ == '__main__':
|
||||
#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:
|
||||
if args.dataset not in dataset_map:
|
||||
raise ValueError('Unrecognized dataset')
|
||||
|
||||
else:
|
||||
train_evaluator = EvaluatorFactory.get_evaluator(dataset_map[args.dataset], model, None, train_iter, args.batch_size, args.gpu)
|
||||
test_evaluator = EvaluatorFactory.get_evaluator(dataset_map[args.dataset], model, None, test_iter, args.batch_size, args.gpu)
|
||||
dev_evaluator = EvaluatorFactory.get_evaluator(dataset_map[args.dataset], model, None, dev_iter, args.batch_size, args.gpu)
|
||||
trainer_config = {
|
||||
'optimizer': optimizer,
|
||||
'batch_size': args.batch_size,
|
||||
@@ -146,17 +138,11 @@ if __name__ == '__main__':
|
||||
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:
|
||||
if args.dataset not in dataset_map:
|
||||
raise ValueError('Unrecognized dataset')
|
||||
else:
|
||||
evaluate_dataset('dev', dataset_map[args.dataset], model, None, dev_iter, args.batch_size, args.gpu)
|
||||
evaluate_dataset('test', dataset_map[args.dataset], model, None, test_iter, args.batch_size, args.gpu)
|
||||
|
||||
# Calculate dev and test metrics
|
||||
for data_loader in [dev_iter, test_iter]:
|
||||
|
||||
+3
-3
@@ -4,7 +4,7 @@ from argparse import ArgumentParser
|
||||
|
||||
|
||||
def get_args():
|
||||
parser = ArgumentParser(description="Kim CNN")
|
||||
parser = ArgumentParser(description="XML 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)
|
||||
@@ -12,12 +12,12 @@ def get_args():
|
||||
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('--dataset', type=str, default='SST-1', choices=['SST-1', 'SST-2', 'Reuters','AAPD'])
|
||||
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('--save_path', type=str, default='xml_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)
|
||||
|
||||
+1
-1
@@ -48,7 +48,7 @@ class XmlCNN(nn.Module):
|
||||
|
||||
|
||||
|
||||
def forward(self, x):
|
||||
def forward(self, x, **kwargs):
|
||||
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)
|
||||
|
||||
Reference in New Issue
Block a user