diff --git a/common/evaluators/reuters_evaluator.py b/common/evaluators/reuters_evaluator.py index 199d93f..dad5063 100644 --- a/common/evaluators/reuters_evaluator.py +++ b/common/evaluators/reuters_evaluator.py @@ -16,12 +16,23 @@ class ReutersEvaluator(Evaluator): self.data_loader.init_epoch() n_dev_correct = 0 total_loss = 0 - + ############ + ## Temp Ave + if hasattr(self.model, 'beta_ema') and self.model.beta_ema > 0: + old_params = self.model.get_params() + self.model.load_ema_params() + ############ for batch_idx, batch in enumerate(self.data_loader): - if self.ignore_lengths: - scores = self.model(batch.text, lengths=batch.text) + if hasattr(self.model, 'TAR') and self.model.TAR: ## TAR Condition + if self.ignore_lengths: + scores, rnn_outs = self.model(batch.text, lengths=batch.text) + else: + scores, rnn_outs = self.model(batch.text[0], lengths=batch.text[1]) else: - scores = self.model(batch.text[0], lengths=batch.text[1]) + if self.ignore_lengths: + scores = self.model(batch.text, lengths=batch.text) + else: + scores = self.model(batch.text[0], lengths=batch.text[1]) scores_rounded = F.sigmoid(scores).round().long() # Using binary accuracy @@ -30,8 +41,14 @@ class ReutersEvaluator(Evaluator): n_dev_correct += 1 total_loss += F.binary_cross_entropy_with_logits(scores, batch.label.float(), size_average=False).item() - + if hasattr(self.model, 'TAR') and self.model.TAR: ### TAR condition + total_loss += (rnn_outs[1:]-rnn_outs[:-1]).pow(2).mean() accuracy = 100. * n_dev_correct / len(self.data_loader.dataset.examples) avg_loss = total_loss / len(self.data_loader.dataset.examples) + ############# + ## Temp Ave + if hasattr(self.model, 'beta_ema') and self.model.beta_ema > 0: + self.model.load_params(old_params) + ############# return [accuracy, avg_loss], ['accuracy', 'cross_entropy_loss'] diff --git a/common/trainers/reuters_trainer.py b/common/trainers/reuters_trainer.py index d2ae689..00e0368 100644 --- a/common/trainers/reuters_trainer.py +++ b/common/trainers/reuters_trainer.py @@ -32,11 +32,16 @@ class ReutersTrainer(Trainer): self.iterations += 1 self.model.train() self.optimizer.zero_grad() - if 'ignore_lengths' in self.config and self.config['ignore_lengths'] == True: - scores = self.model(batch.text, lengths=batch.text) + if hasattr(self.model, 'TAR') and self.model.TAR: + if 'ignore_lengths' in self.config and self.config['ignore_lengths'] == True: + scores, rnn_outs = self.model(batch.text, lengths=batch.text) + else: + scores, rnn_outs = self.model(batch.text[0], lengths=batch.text[1]) else: - scores = self.model(batch.text[0], lengths=batch.text[1]) - + if 'ignore_lengths' in self.config and self.config['ignore_lengths'] == True: + scores = self.model(batch.text, lengths=batch.text) + else: + scores = self.model(batch.text[0], lengths=batch.text[1]) # Using binary accuracy for tensor1, tensor2 in zip(F.sigmoid(scores).round().long(), batch.label): if np.array_equal(tensor1, tensor2): @@ -44,9 +49,16 @@ class ReutersTrainer(Trainer): n_total += batch.batch_size train_acc = 100. * n_correct / n_total 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.backward() self.optimizer.step() + ############# + ## Temp Ave + if hasattr(self.model, 'beta_ema') and self.model.beta_ema > 0: + self.model.update_ema() + ############# # Evaluate performance on validation set if self.iterations % self.dev_log_interval == 1: diff --git a/lstm_baseline/__main__.py b/lstm_baseline/__main__.py index 4071249..2b2eea7 100644 --- a/lstm_baseline/__main__.py +++ b/lstm_baseline/__main__.py @@ -171,7 +171,7 @@ if __name__ == '__main__': predicted_labels = list() target_labels = list() for batch_idx, batch in enumerate(data_loader): - scores_rounded = F.sigmoid(model(batch.text)).round().long() + scores_rounded = F.sigmoid(model(batch.text[0])).round().long() predicted_labels.extend(scores_rounded.cpu().detach().numpy()) target_labels.extend(batch.label.cpu().detach().numpy()) predicted_labels = np.array(predicted_labels) diff --git a/lstm_regularization/README.md b/lstm_regularization/README.md new file mode 100644 index 0000000..8aa627d --- /dev/null +++ b/lstm_regularization/README.md @@ -0,0 +1,42 @@ +# lstm_baseline with Regularization + +Implementation of a standard LSTM using PyTorch and Torchtext for text classification baseline measurements with Regularization. + +## 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 lstm_baseline --mode static +``` + +## Dataset + +We experiment the model on the following datasets. + +- Reuters dataset - ModApte splits + +## Settings + +Adam is used for training with an option of temporal averaging. + +## TODO +- Support ONNX export. Currently throws a ONNX export failed (Couldn't export Python operator forward_flattened_wrapper) exception. +- Add dataset results with different hyperparameters +- Parameters tuning + +## Regularization Module + +- Regularization methods like Embedding dropout, Weight Dropped LSTM and Temporal Activation Regularization are implemented. +- Temporal Averaging is also an additional module + +## Acknowledgement +- The additional modules have been heavily inspired by two open source repositories: + - https://github.com/salesforce/awd-lstm-lm.git + - https://github.com/AMLab-Amsterdam/L0_regularization.git diff --git a/lstm_regularization/__init__.py b/lstm_regularization/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/lstm_regularization/__main__.py b/lstm_regularization/__main__.py new file mode 100644 index 0000000..abea771 --- /dev/null +++ b/lstm_regularization/__main__.py @@ -0,0 +1,171 @@ +from copy import deepcopy +import logging +import random + +import numpy as np +import torch +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 datasets.aapd import AAPD +from lstm_regularization.args import get_args +from lstm_regularization.model import LSTMBaseline + +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() + + 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 + 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 = LSTMBaseline(config) + if args.cuda: + model.cuda() + print('Shift model to GPU') + + parameter = filter(lambda p: p.requires_grad, model.parameters()) + optimizer = torch.optim.Adam(parameter, lr=args.lr, weight_decay=args.weight_decay) + + 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, + '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 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 + if model.beta_ema > 0: + old_params = model.get_params() + model.load_ema_params() + + for data_loader in [dev_iter, test_iter]: + predicted_labels = list() + target_labels = list() + for batch_idx, batch in enumerate(data_loader): + if model.TAR: + scores_rounded = F.sigmoid(model(batch.text[0])[0]).round().long() + else: + scores_rounded = F.sigmoid(model(batch.text[0])).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 model.beta_ema > 0: + model.load_params(old_params) diff --git a/lstm_regularization/args.py b/lstm_regularization/args.py new file mode 100644 index 0000000..3e8f638 --- /dev/null +++ b/lstm_regularization/args.py @@ -0,0 +1,41 @@ +import os + +from argparse import ArgumentParser + + +def get_args(): + parser = ArgumentParser(description="Baseline LSTM for text classification with Regularization") + 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=50) + parser.add_argument('--batch_size', type=int, default=1024) + parser.add_argument('--bidirectional', action='store_true'), + parser.add_argument('--bottleneck_layer', action='store_true'), + parser.add_argument('--num_layers', type=int, default=2) + parser.add_argument('--hidden_dim', type=int, default=256) + parser.add_argument('--mode', type=str, default='static', choices=['rand', 'static', 'non-static']) + parser.add_argument('--lr', type=float, default=0.001) + parser.add_argument('--seed', type=int, default=3435) + parser.add_argument('--dataset', type=str, default='Reuters', 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='lstm_regularization/saves') + 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='GoogleNews-vectors-negative300.txt') + parser.add_argument('--trained_model', type=str, default="") + parser.add_argument('--TAR', action='store_true') + 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") + parser.add_argument('--embed_droprate', type=float, default=0.0, help="for embedded droupout") + args = parser.parse_args() + return args diff --git a/lstm_regularization/embed_regularize.py b/lstm_regularization/embed_regularize.py new file mode 100644 index 0000000..0e47cdf --- /dev/null +++ b/lstm_regularization/embed_regularize.py @@ -0,0 +1,53 @@ +""" +BSD 3-Clause License + +Copyright (c) 2017, +All rights reserved. + +Redistribution and use in source and binary forms, with or without +modification, are permitted provided that the following conditions are met: + +* Redistributions of source code must retain the above copyright notice, this + list of conditions and the following disclaimer. + +* Redistributions in binary form must reproduce the above copyright notice, + this list of conditions and the following disclaimer in the documentation + and/or other materials provided with the distribution. + +* Neither the name of the copyright holder nor the names of its + contributors may be used to endorse or promote products derived from + this software without specific prior written permission. + +THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE +DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE +FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR +SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER +CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, +OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. +""" +import numpy as np + +import torch + +def embedded_dropout(embed, words, dropout=0.1, scale=None): + if dropout: + mask = embed.weight.data.new().resize_((embed.weight.size(0), 1)).bernoulli_(1 - dropout).expand_as(embed.weight) / (1 - dropout) + masked_embed_weight = mask * embed.weight + else: + masked_embed_weight = embed.weight + if scale: + masked_embed_weight = scale.expand_as(masked_embed_weight) * masked_embed_weight + + padding_idx = embed.padding_idx + if padding_idx is None: + padding_idx = -1 + + X = torch.nn.functional.embedding(words, masked_embed_weight, + padding_idx, embed.max_norm, embed.norm_type, + embed.scale_grad_by_freq, embed.sparse + ) + return X diff --git a/lstm_regularization/locked_dropout.py b/lstm_regularization/locked_dropout.py new file mode 100644 index 0000000..65a2416 --- /dev/null +++ b/lstm_regularization/locked_dropout.py @@ -0,0 +1,15 @@ +import torch +import torch.nn as nn +from torch.autograd import Variable + +class LockedDropout(nn.Module): + def __init__(self): + super().__init__() + + def forward(self, x, dropout=0.5): + if not self.training or not dropout: + return x + m = x.data.new(1, x.size(1), x.size(2)).bernoulli_(1 - dropout) + mask = m / (1 - dropout) + mask = mask.expand_as(x) + return mask * x diff --git a/lstm_regularization/model.py b/lstm_regularization/model.py new file mode 100644 index 0000000..eadbe07 --- /dev/null +++ b/lstm_regularization/model.py @@ -0,0 +1,109 @@ +from copy import deepcopy +import torch +import torch.nn as nn + +import torch.nn.functional as F + +from lstm_regularization.weight_drop import WeightDrop +from lstm_regularization.embed_regularize import embedded_dropout + + +class LSTMBaseline(nn.Module): + def __init__(self, config): + super(LSTMBaseline, self).__init__() + dataset = config.dataset + target_class = config.target_class + self.is_bidirectional = config.bidirectional + self.has_bottleneck_layer = config.bottleneck_layer + self.mode = config.mode + self.TAR = config.TAR + self.beta_ema = config.beta_ema ## Temporal averaging + self.wdrop = config.wdrop ## WEight dropping + self.embed_droprate = config.embed_droprate ## Embedding dropout + + input_channel = 1 + if config.mode == 'rand': + rand_embed_init = torch.Tensor(config.words_num, config.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) + else: + print("Unsupported Mode") + exit() + + self.lstm = nn.LSTM(config.words_dim, config.hidden_dim, dropout=config.dropout, num_layers=config.num_layers, + bidirectional=self.is_bidirectional, batch_first=True) + ## Wdrop + if self.wdrop: + self.lstm = WeightDrop(self.lstm, ['weight_hh_l0'], dropout=self.wdrop) + self.dropout = nn.Dropout(config.dropout) + if self.has_bottleneck_layer: + if self.is_bidirectional: + self.fc1 = nn.Linear(2 * config.hidden_dim, config.hidden_dim) # Hidden Bottleneck Layer + self.fc2 = nn.Linear(config.hidden_dim, target_class) + else: + self.fc1 = nn.Linear(config.hidden_dim, config.hidden_dim//2) # Hidden Bottleneck Layer + self.fc2 = nn.Linear(config.hidden_dim//2, target_class) + else: + if self.is_bidirectional: + self.fc1 = nn.Linear(2 * config.hidden_dim, target_class) + else: + self.fc1 = nn.Linear(config.hidden_dim, target_class) + + if self.beta_ema>0: + self.avg_param = deepcopy(list(p.data for p in self.parameters())) + if torch.cuda.is_available(): + self.avg_param = [a.cuda() for a in self.avg_param] + self.steps_ema = 0. + + def forward(self, x, lengths=None): + if self.mode == 'rand': + x = embedded_dropout(self.embed, x, dropout=self.embed_droprate if self.training else 0) if self.embed_droprate else self.embed(x) + elif self.mode == 'static': + x = embedded_dropout(self.static_embed, x, dropout=self.embed_droprate if self.training else 0) if self.embed_droprate else self.static_embed(x) + elif self.mode == 'non-static': + x = embedded_dropout(self.non_static_embed, x, dropout=self.embed_droprate if self.training else 0) if self.embed_droprate else self.non_static_embed(x) + else: + print("Unsupported Mode") + exit() + if lengths is not None: + x = torch.nn.utils.rnn.pack_padded_sequence(x, lengths, batch_first=True) + rnn_outs, _ = self.lstm(x) + rnn_outs_temp = rnn_outs + #rnn_outs,_ = torch.nn.utils.rnn.pad_packed_sequence(rnn_outs, batch_first = True) + + if lengths is not None: + rnn_outs,_ = torch.nn.utils.rnn.pad_packed_sequence(rnn_outs, batch_first=True) + rnn_outs_temp, _ = torch.nn.utils.rnn.pad_packed_sequence(rnn_outs_temp, batch_first=True) + + x = F.relu(torch.transpose(rnn_outs_temp, 1, 2)) + x = F.max_pool1d(x, x.size(2)).squeeze(2) + x = self.dropout(x) + if self.has_bottleneck_layer: + x = F.relu(self.fc1(x)) + if self.TAR: + return self.fc2(x), rnn_outs.permute(1,0,2) + return self.fc2(x) + else: + if self.TAR: + return self.fc1(x), rnn_outs.permute(1,0,2) + return self.fc1(x) + + def update_ema(self): + self.steps_ema += 1 + for p, avg_p in zip(self.parameters(), self.avg_param): + avg_p.mul_(self.beta_ema).add_((1-self.beta_ema)*p.data) + + def load_ema_params(self): + for p, avg_p in zip(self.parameters(), self.avg_param): + p.data.copy_(avg_p/(1-self.beta_ema**self.steps_ema)) + + def load_params(self, params): + for p,avg_p in zip(self.parameters(), params): + p.data.copy_(avg_p) + + def get_params(self): + params = deepcopy(list(p.data for p in self.parameters())) + return params diff --git a/lstm_regularization/weight_drop.py b/lstm_regularization/weight_drop.py new file mode 100644 index 0000000..d1c473d --- /dev/null +++ b/lstm_regularization/weight_drop.py @@ -0,0 +1,78 @@ +""" +BSD 3-Clause License + +Copyright (c) 2017, +All rights reserved. + +Redistribution and use in source and binary forms, with or without +modification, are permitted provided that the following conditions are met: + +* Redistributions of source code must retain the above copyright notice, this + list of conditions and the following disclaimer. + +* Redistributions in binary form must reproduce the above copyright notice, + this list of conditions and the following disclaimer in the documentation + and/or other materials provided with the distribution. + +* Neither the name of the copyright holder nor the names of its + contributors may be used to endorse or promote products derived from + this software without specific prior written permission. + +THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE +DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE +FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR +SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER +CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, +OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. +""" +import torch +from torch.nn import Parameter +from functools import wraps + +class WeightDrop(torch.nn.Module): + def __init__(self, module, weights, dropout=0, variational=False): + super(WeightDrop, self).__init__() + self.module = module + self.weights = weights + self.dropout = dropout + self.variational = variational + self._setup() + + def widget_demagnetizer_y2k_edition(*args, **kwargs): + # We need to replace flatten_parameters with a nothing function + # It must be a function rather than a lambda as otherwise pickling explodes + # We can't write boring code though, so ... WIDGET DEMAGNETIZER Y2K EDITION! + + return + + def _setup(self): + # Terrible temporary solution to an issue regarding compacting weights re: CUDNN RNN + if issubclass(type(self.module), torch.nn.RNNBase): + self.module.flatten_parameters = self.widget_demagnetizer_y2k_edition + + for name_w in self.weights: + print('Applying weight drop of {} to {}'.format(self.dropout, name_w)) + w = getattr(self.module, name_w) + del self.module._parameters[name_w] + self.module.register_parameter(name_w + '_raw', Parameter(w.data)) + + def _setweights(self): + for name_w in self.weights: + raw_w = getattr(self.module, name_w + '_raw') + w = None + if self.variational: + mask = torch.autograd.Variable(torch.ones(raw_w.size(0), 1)) + if raw_w.is_cuda: mask = mask.cuda() + mask = torch.nn.functional.dropout(mask, p=self.dropout, training=True) + w = mask.expand_as(raw_w) * raw_w + else: + w = torch.nn.functional.dropout(raw_w, p=self.dropout, training=self.training) + setattr(self.module, name_w, w) + + def forward(self, *args): + self._setweights() + return self.module.forward(*args)