mirror of
https://github.com/wassname/Castor.git
synced 2026-09-09 11:13:20 +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:
@@ -14,11 +14,14 @@ class ReutersEvaluator(Evaluator):
|
||||
total_loss = 0
|
||||
|
||||
for batch_idx, batch in enumerate(self.data_loader):
|
||||
scores = self.model(batch.text)
|
||||
scores = self.model(batch.text[0], lengths=batch.text[1])
|
||||
scores_rounded = F.sigmoid(scores).round().long()
|
||||
|
||||
# Using binary accuracy
|
||||
for tensor1, tensor2 in zip(F.sigmoid(scores).round().long(), batch.label):
|
||||
for tensor1, tensor2 in zip(scores_rounded, batch.label):
|
||||
if np.array_equal(tensor1, tensor2):
|
||||
n_dev_correct += 1
|
||||
|
||||
total_loss += F.binary_cross_entropy_with_logits(scores, batch.label.float(), size_average=False).item()
|
||||
|
||||
accuracy = 100. * n_dev_correct / len(self.data_loader.dataset.examples)
|
||||
|
||||
Reference in New Issue
Block a user