From d326e575c671ce670d2904622b298e72dbbaf0a8 Mon Sep 17 00:00:00 2001 From: Ashutosh-Adhikari Date: Fri, 25 Jan 2019 16:39:35 -0500 Subject: [PATCH] Add TAR and AR (#172) * Add TAR and AR --- common/trainers/reuters_trainer.py | 4 +++- lstm_regularization/args.py | 3 ++- lstm_regularization/model.py | 5 +++-- 3 files changed, 8 insertions(+), 4 deletions(-) diff --git a/common/trainers/reuters_trainer.py b/common/trainers/reuters_trainer.py index 69b8e3c..483de7f 100644 --- a/common/trainers/reuters_trainer.py +++ b/common/trainers/reuters_trainer.py @@ -58,7 +58,9 @@ class ReutersTrainer(Trainer): loss = F.binary_cross_entropy_with_logits(scores, batch.label.float()) if hasattr(self.model, 'TAR') and self.model.TAR: - loss = loss + (rnn_outs[1:] - rnn_outs[:-1]).pow(2).mean() + loss = loss + self.model.TAR*(rnn_outs[1:] - rnn_outs[:-1]).pow(2).mean() + if hasattr(self.model, 'AR') and self.model.AR: + loss = loss + self.model.AR*(rnn_outs[:]).pow(2).mean() n_total += batch.batch_size train_acc = 100. * n_correct / n_total diff --git a/lstm_regularization/args.py b/lstm_regularization/args.py index 3e823e2..d264ab5 100644 --- a/lstm_regularization/args.py +++ b/lstm_regularization/args.py @@ -33,7 +33,8 @@ def get_args(): 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('--TAR', action='store_true') + parser.add_argument('--TAR', type=float, default=0.0, help="Hyperparameter for Temporal Activation Regularization") + parser.add_argument('--AR', type=float, default=0.0, help="Hyperparameter for Activation Regularization") parser.add_argument('--weight_decay', type=float, default=0) parser.add_argument('--beta_ema', type=float, default = 0, help="for temporal averaging") parser.add_argument('--wdrop', type=float, default=0.0, help="for weight-drop") diff --git a/lstm_regularization/model.py b/lstm_regularization/model.py index e19879b..2117091 100644 --- a/lstm_regularization/model.py +++ b/lstm_regularization/model.py @@ -17,6 +17,7 @@ class LSTMBaseline(nn.Module): self.has_bottleneck_layer = config.bottleneck_layer self.mode = config.mode self.TAR = config.TAR + self.AR = config.AR self.beta_ema = config.beta_ema ## Temporal averaging self.wdrop = config.wdrop ## Weight dropping self.embed_droprate = config.embed_droprate ## Embedding dropout @@ -84,11 +85,11 @@ class LSTMBaseline(nn.Module): if self.has_bottleneck_layer: x = F.relu(self.fc1(x)) # x = self.dropout(x) - if self.TAR: + if self.TAR or self.AR: return self.fc2(x), rnn_outs.permute(1,0,2) return self.fc2(x) else: - if self.TAR: + if self.TAR or self.AR: return self.fc1(x), rnn_outs.permute(1,0,2) return self.fc1(x)