Loading fairseq/criterions/adaptive_loss.py +11 −6 Original line number Diff line number Diff line Loading @@ -42,11 +42,14 @@ class AdaptiveLoss(FairseqCriterion): adaptive_softmax = model.decoder.adaptive_softmax net_output = model(**sample['net_input']) target = model.get_targets(sample, net_output).view(-1) orig_target = model.get_targets(sample, net_output) bsz = target.size(0) nsentences = orig_target.size(0) orig_target = orig_target.view(-1) logits, target = adaptive_softmax(net_output[0], target) bsz = orig_target.size(0) logits, target = adaptive_softmax(net_output[0], orig_target) assert len(target) == len(logits) loss = net_output[0].new(1 if reduce else bsz).zero_() Loading @@ -57,11 +60,13 @@ class AdaptiveLoss(FairseqCriterion): loss += F.cross_entropy(logits[i], target[i], size_average=False, ignore_index=self.padding_idx, reduce=reduce) sample_size = sample['target'].size(0) if self.args.sentence_avg else sample['ntokens'] orig = utils.strip_pad(orig_target, self.padding_idx) ntokens = orig.numel() sample_size = sample['target'].size(0) if self.args.sentence_avg else ntokens logging_output = { 'loss': utils.item(loss.data) if reduce else loss.data, 'ntokens': sample['ntokens'], 'nsentences': sample['target'].size(0), 'ntokens': ntokens, 'nsentences': nsentences, 'sample_size': sample_size, } return loss, sample_size, logging_output Loading fairseq/data/__init__.py +1 −1 Original line number Diff line number Diff line Loading @@ -5,7 +5,7 @@ # the root directory of this source tree. An additional grant of patent rights # can be found in the PATENTS file in the same directory. from .dictionary import Dictionary from .dictionary import Dictionary, TruncatedDictionary from .fairseq_dataset import FairseqDataset from .indexed_dataset import IndexedDataset, IndexedInMemoryDataset, IndexedRawTextDataset from .language_pair_dataset import LanguagePairDataset Loading fairseq/data/dictionary.py +17 −0 Original line number Diff line number Diff line Loading @@ -199,3 +199,20 @@ class Dictionary(object): t = torch.Tensor(length).uniform_(self.nspecial + 1, len(self)).long() t[-1] = self.eos() return t class TruncatedDictionary(object): def __init__(self, wrapped_dict, length): self.__class__ = type(dict.__class__.__name__, (self.__class__, dict.__class__), {}) self.__dict__ = dict.__dict__ self.wrapped_dict = wrapped_dict self.length = min(len(self.wrapped_dict), length) def __len__(self): return self.length def __getitem__(self, i): if i < self.length: return self.wrapped_dict[i] return self.wrapped_dict.unk() fairseq/data/language_pair_dataset.py +1 −0 Original line number Diff line number Diff line Loading @@ -61,6 +61,7 @@ def collate( 'src_lengths': src_lengths, }, 'target': target, 'nsentences': samples[0]['source'].size(0), } if prev_output_tokens is not None: batch['net_input']['prev_output_tokens'] = prev_output_tokens Loading fairseq/data/monolingual_dataset.py +78 −8 Original line number Diff line number Diff line Loading @@ -9,27 +9,39 @@ import numpy as np import torch from . import data_utils, FairseqDataset from typing import List def collate(samples, pad_idx, eos_idx): if len(samples) == 0: return {} def merge(key): def merge(key, is_list=False): if is_list: res = [] for i in range(len(samples[0][key])): res.append(data_utils.collate_tokens( [s[key][i] for s in samples], pad_idx, eos_idx, left_pad=False, )) return res else: return data_utils.collate_tokens( [s[key] for s in samples], pad_idx, eos_idx, left_pad=False, ) is_target_list = isinstance(samples[0]['target'], list) return { 'id': torch.LongTensor([s['id'] for s in samples]), 'ntokens': sum(len(s['target']) for s in samples), 'ntokens': sum(len(s['source']) for s in samples), 'net_input': { 'src_tokens': merge('source'), 'src_lengths': torch.LongTensor([ s['source'].numel() for s in samples ]), }, 'target': merge('target'), 'target': merge('target', is_target_list), 'nsentences': samples[0]['source'].size(0), } Loading @@ -45,19 +57,75 @@ class MonolingualDataset(FairseqDataset): Default: ``True`` """ def __init__(self, dataset, sizes, vocab, shuffle=True): def __init__(self, dataset, sizes, src_vocab, tgt_vocab, add_eos_for_other_targets, shuffle, targets=None): self.dataset = dataset self.sizes = np.array(sizes) self.vocab = vocab self.vocab = src_vocab self.tgt_vocab = tgt_vocab self.add_eos_for_other_targets = add_eos_for_other_targets self.shuffle = shuffle assert targets is None or all( t in {'self', 'future', 'past'} for t in targets), "targets must be none or one of 'self', 'future', 'past'" if targets is not None and len(targets) == 0: targets = None self.targets = targets def __getitem__(self, index): source, target = self.dataset[index] source, future_target, past_target = self.dataset[index] source, target = self._make_source_target(source, future_target, past_target) return {'id': index, 'source': source, 'target': target} def __len__(self): return len(self.dataset) def _make_source_target(self, source, future_target, past_target): if self.targets is not None: target = [] if self.add_eos_for_other_targets and (('self' in self.targets) or ('past' in self.targets)) \ and source[-1] != self.vocab.eos(): # append eos at the end of source source = torch.cat([source, source.new([self.vocab.eos()])]) if 'future' in self.targets: future_target = torch.cat([future_target, future_target.new([self.vocab.pad()])]) if 'past' in self.targets: # first token is before the start of sentence which is only used in "none" break mode when # add_eos_for_other_targets is False past_target = torch.cat([past_target.new([self.vocab.pad()]), past_target[1:], source[-2, None]]) for t in self.targets: if t == 'self': target.append(source) elif t == 'future': target.append(future_target) elif t == 'past': target.append(past_target) else: raise Exception('invalid target ' + t) if len(target) == 1: target = target[0] else: target = future_target return source, self._filter_vocab(target) def _filter_vocab(self, target): if len(self.tgt_vocab) != len(self.vocab): def _filter(target): mask = target.ge(len(self.tgt_vocab)) if mask.any(): target[mask] = self.tgt_vocab.unk() return target if isinstance(target, list): return [_filter(t) for t in target] return _filter(target) return target def collater(self, samples): """Merge a list of samples to form a mini-batch. Loading Loading @@ -86,8 +154,10 @@ class MonolingualDataset(FairseqDataset): if isinstance(max_positions, float) or isinstance(max_positions, int): tgt_len = min(tgt_len, max_positions) bsz = num_tokens // tgt_len target = self.vocab.dummy_sentence(tgt_len + 1) source, target = target[:-1], target[1:] target = self.vocab.dummy_sentence(tgt_len + 2) source, past_target, future_target = target[1:-1], target[2:], target[:-2] source, target = self._make_source_target(source, past_target, future_target) return self.collater([ {'id': i, 'source': source, 'target': target} for i in range(bsz) Loading Loading
fairseq/criterions/adaptive_loss.py +11 −6 Original line number Diff line number Diff line Loading @@ -42,11 +42,14 @@ class AdaptiveLoss(FairseqCriterion): adaptive_softmax = model.decoder.adaptive_softmax net_output = model(**sample['net_input']) target = model.get_targets(sample, net_output).view(-1) orig_target = model.get_targets(sample, net_output) bsz = target.size(0) nsentences = orig_target.size(0) orig_target = orig_target.view(-1) logits, target = adaptive_softmax(net_output[0], target) bsz = orig_target.size(0) logits, target = adaptive_softmax(net_output[0], orig_target) assert len(target) == len(logits) loss = net_output[0].new(1 if reduce else bsz).zero_() Loading @@ -57,11 +60,13 @@ class AdaptiveLoss(FairseqCriterion): loss += F.cross_entropy(logits[i], target[i], size_average=False, ignore_index=self.padding_idx, reduce=reduce) sample_size = sample['target'].size(0) if self.args.sentence_avg else sample['ntokens'] orig = utils.strip_pad(orig_target, self.padding_idx) ntokens = orig.numel() sample_size = sample['target'].size(0) if self.args.sentence_avg else ntokens logging_output = { 'loss': utils.item(loss.data) if reduce else loss.data, 'ntokens': sample['ntokens'], 'nsentences': sample['target'].size(0), 'ntokens': ntokens, 'nsentences': nsentences, 'sample_size': sample_size, } return loss, sample_size, logging_output Loading
fairseq/data/__init__.py +1 −1 Original line number Diff line number Diff line Loading @@ -5,7 +5,7 @@ # the root directory of this source tree. An additional grant of patent rights # can be found in the PATENTS file in the same directory. from .dictionary import Dictionary from .dictionary import Dictionary, TruncatedDictionary from .fairseq_dataset import FairseqDataset from .indexed_dataset import IndexedDataset, IndexedInMemoryDataset, IndexedRawTextDataset from .language_pair_dataset import LanguagePairDataset Loading
fairseq/data/dictionary.py +17 −0 Original line number Diff line number Diff line Loading @@ -199,3 +199,20 @@ class Dictionary(object): t = torch.Tensor(length).uniform_(self.nspecial + 1, len(self)).long() t[-1] = self.eos() return t class TruncatedDictionary(object): def __init__(self, wrapped_dict, length): self.__class__ = type(dict.__class__.__name__, (self.__class__, dict.__class__), {}) self.__dict__ = dict.__dict__ self.wrapped_dict = wrapped_dict self.length = min(len(self.wrapped_dict), length) def __len__(self): return self.length def __getitem__(self, i): if i < self.length: return self.wrapped_dict[i] return self.wrapped_dict.unk()
fairseq/data/language_pair_dataset.py +1 −0 Original line number Diff line number Diff line Loading @@ -61,6 +61,7 @@ def collate( 'src_lengths': src_lengths, }, 'target': target, 'nsentences': samples[0]['source'].size(0), } if prev_output_tokens is not None: batch['net_input']['prev_output_tokens'] = prev_output_tokens Loading
fairseq/data/monolingual_dataset.py +78 −8 Original line number Diff line number Diff line Loading @@ -9,27 +9,39 @@ import numpy as np import torch from . import data_utils, FairseqDataset from typing import List def collate(samples, pad_idx, eos_idx): if len(samples) == 0: return {} def merge(key): def merge(key, is_list=False): if is_list: res = [] for i in range(len(samples[0][key])): res.append(data_utils.collate_tokens( [s[key][i] for s in samples], pad_idx, eos_idx, left_pad=False, )) return res else: return data_utils.collate_tokens( [s[key] for s in samples], pad_idx, eos_idx, left_pad=False, ) is_target_list = isinstance(samples[0]['target'], list) return { 'id': torch.LongTensor([s['id'] for s in samples]), 'ntokens': sum(len(s['target']) for s in samples), 'ntokens': sum(len(s['source']) for s in samples), 'net_input': { 'src_tokens': merge('source'), 'src_lengths': torch.LongTensor([ s['source'].numel() for s in samples ]), }, 'target': merge('target'), 'target': merge('target', is_target_list), 'nsentences': samples[0]['source'].size(0), } Loading @@ -45,19 +57,75 @@ class MonolingualDataset(FairseqDataset): Default: ``True`` """ def __init__(self, dataset, sizes, vocab, shuffle=True): def __init__(self, dataset, sizes, src_vocab, tgt_vocab, add_eos_for_other_targets, shuffle, targets=None): self.dataset = dataset self.sizes = np.array(sizes) self.vocab = vocab self.vocab = src_vocab self.tgt_vocab = tgt_vocab self.add_eos_for_other_targets = add_eos_for_other_targets self.shuffle = shuffle assert targets is None or all( t in {'self', 'future', 'past'} for t in targets), "targets must be none or one of 'self', 'future', 'past'" if targets is not None and len(targets) == 0: targets = None self.targets = targets def __getitem__(self, index): source, target = self.dataset[index] source, future_target, past_target = self.dataset[index] source, target = self._make_source_target(source, future_target, past_target) return {'id': index, 'source': source, 'target': target} def __len__(self): return len(self.dataset) def _make_source_target(self, source, future_target, past_target): if self.targets is not None: target = [] if self.add_eos_for_other_targets and (('self' in self.targets) or ('past' in self.targets)) \ and source[-1] != self.vocab.eos(): # append eos at the end of source source = torch.cat([source, source.new([self.vocab.eos()])]) if 'future' in self.targets: future_target = torch.cat([future_target, future_target.new([self.vocab.pad()])]) if 'past' in self.targets: # first token is before the start of sentence which is only used in "none" break mode when # add_eos_for_other_targets is False past_target = torch.cat([past_target.new([self.vocab.pad()]), past_target[1:], source[-2, None]]) for t in self.targets: if t == 'self': target.append(source) elif t == 'future': target.append(future_target) elif t == 'past': target.append(past_target) else: raise Exception('invalid target ' + t) if len(target) == 1: target = target[0] else: target = future_target return source, self._filter_vocab(target) def _filter_vocab(self, target): if len(self.tgt_vocab) != len(self.vocab): def _filter(target): mask = target.ge(len(self.tgt_vocab)) if mask.any(): target[mask] = self.tgt_vocab.unk() return target if isinstance(target, list): return [_filter(t) for t in target] return _filter(target) return target def collater(self, samples): """Merge a list of samples to form a mini-batch. Loading Loading @@ -86,8 +154,10 @@ class MonolingualDataset(FairseqDataset): if isinstance(max_positions, float) or isinstance(max_positions, int): tgt_len = min(tgt_len, max_positions) bsz = num_tokens // tgt_len target = self.vocab.dummy_sentence(tgt_len + 1) source, target = target[:-1], target[1:] target = self.vocab.dummy_sentence(tgt_len + 2) source, past_target, future_target = target[1:-1], target[2:], target[:-2] source, target = self._make_source_target(source, past_target, future_target) return self.collater([ {'id': i, 'source': source, 'target': target} for i in range(bsz) Loading