Commit bd4db8fb authored by Myle Ott's avatar Myle Ott
Browse files

Misc changes for pytorch-translate

parent c6fe9fc5
Loading
Loading
Loading
Loading
+4 −4
Changes for fairseq/data/dictionary.py: 4 added lines, 4 removed lines.
Original line number Diff line number Diff line
@@ -106,7 +106,7 @@ class Dictionary(object):
                multiple of 8, which is important on some hardware (e.g., Nvidia
                Tensor Cores).
        """
        if nwords == -1:
        if nwords <= 0:
            nwords = len(self)

        new_indices = dict(zip(self.symbols[:self.nspecial], range(self.nspecial)))
@@ -133,7 +133,7 @@ class Dictionary(object):
                i += 1
                threshold_nwords += 1

        assert min(new_count[self.nspecial:]) >= threshold
        assert len(new_count) == self.nspecial or min(new_count[self.nspecial:]) >= threshold
        assert len(new_symbols) % padding_factor == 0
        assert len(new_symbols) == len(new_indices)

@@ -187,12 +187,12 @@ class Dictionary(object):
            d.count.append(count)
        return d

    def save(self, f, threshold=3, nwords=-1):
    def save(self, f):
        """Stores dictionary into a text file"""
        if isinstance(f, str):
            os.makedirs(os.path.dirname(f), exist_ok=True)
            with open(f, 'w', encoding='utf-8') as fd:
                return self.save(fd, threshold, nwords)
                return self.save(fd)
        for symbol, count in zip(self.symbols[self.nspecial:], self.count[self.nspecial:]):
            print('{} {}'.format(symbol, count), file=f)

+7 −2
Changes for fairseq/data/indexed_dataset.py: 7 added lines, 2 removed lines.
Original line number Diff line number Diff line
@@ -52,8 +52,9 @@ def data_file_path(prefix_path):
class IndexedDataset(torch.utils.data.Dataset):
    """Loader for TorchNet IndexedDataset"""

    def __init__(self, path):
    def __init__(self, path, fix_lua_indexing=False):
        super().__init__()
        self.fix_lua_indexing = fix_lua_indexing
        with open(index_file_path(path), 'rb') as f:
            magic = f.read(8)
            assert magic == b'TNTIDX\x00\x00'
@@ -83,7 +84,10 @@ class IndexedDataset(torch.utils.data.Dataset):
        a = np.empty(tensor_size, dtype=self.dtype)
        self.data_file.seek(self.data_offsets[i] * self.element_size)
        self.data_file.readinto(a)
        return torch.from_numpy(a).long() - 1  # subtract 1 for 0-based indexing
        item = torch.from_numpy(a).long()
        if self.fix_lua_indexing:
            item -= 1  # subtract 1 for 0-based indexing
        return item

    def __len__(self):
        return self.size
@@ -104,6 +108,7 @@ class IndexedInMemoryDataset(IndexedDataset):
        self.buffer = np.empty(self.data_offsets[-1], dtype=self.dtype)
        self.data_file.readinto(self.buffer)
        self.data_file.close()
        if self.fix_lua_indexing:
            self.buffer -= 1  # subtract 1 for 0-based indexing

    def __del__(self):
+1 −1
Changes for fairseq/fp16_trainer.py: 1 added line, 1 removed line.
Original line number Diff line number Diff line
@@ -73,7 +73,7 @@ class FP16Trainer(Trainer):
        self.fp32_params.grad = self.fp32_params.data.new(total_param_size)

        # create optimizer using the copied FP32 params
        self.optimizer = optim.build_optimizer(self.args, [self.fp32_params])
        self._optimizer = optim.build_optimizer(self.args, [self.fp32_params])
        self.lr_scheduler = lr_scheduler.build_lr_scheduler(self.args, self.optimizer)

    def save_checkpoint(self, filename, extra_state):
+3 −0
Changes for fairseq/optim/lr_scheduler/fixed_schedule.py: 3 added lines, 0 removed lines.
Original line number Diff line number Diff line
@@ -15,6 +15,9 @@ class FixedSchedule(FairseqLRScheduler):
    def __init__(self, args, optimizer):
        super().__init__(args, optimizer)

        # set defaults
        args.warmup_updates = getattr(args, 'warmup_updates', 0)

        self.lr = args.lr[0]
        if args.warmup_updates > 0:
            self.warmup_factor = 1. / args.warmup_updates
+1 −1
Changes for fairseq/tasks/language_modeling.py: 1 added line, 1 removed line.
Original line number Diff line number Diff line
@@ -50,7 +50,7 @@ class LanguageModelingTask(FairseqTask):
            ds = IndexedRawTextDataset(path, self.dictionary)
            tokens = ds.tokens_list
        elif not self.args.raw_text and IndexedInMemoryDataset.exists(path):
            ds = IndexedInMemoryDataset(path)
            ds = IndexedInMemoryDataset(path, fix_lua_indexing=True)
            tokens = ds.buffer
        else:
            raise FileNotFoundError('Dataset not found: {} ({})'.format(split, self.args.data))
Loading