Unverified Commit d3795d6c authored by Myle Ott's avatar Myle Ott Committed by GitHub
Browse files

Merge internal changes (#136)

Changes:
- 7d19e36: Add `--sampling` flag to generate.py to sample instead of doing beam search
- c777340: Add `scripts/average_checkpoints.py` to average multiple checkpoints into a combined model
- 3ea882c: Add `--max-update` option to train.py to stop training after a given number of updates
- small bugfixes for distributed training, LSTM, inverse square root LR scheduler
parent 48836525
Loading
Loading
Loading
Loading
+5 −2
Changes for fairseq/criterions/cross_entropy.py: 5 added lines, 2 removed lines.
Original line number Diff line number Diff line
@@ -8,9 +8,11 @@
import math
import torch.nn.functional as F

from . import FairseqCriterion, register_criterion
from fairseq import utils

from . import FairseqCriterion, register_criterion


@register_criterion('cross_entropy')
class CrossEntropyCriterion(FairseqCriterion):

@@ -28,7 +30,7 @@ class CrossEntropyCriterion(FairseqCriterion):
        net_output = model(**sample['net_input'])
        lprobs = model.get_normalized_probs(net_output, log_probs=True)
        lprobs = lprobs.view(-1, lprobs.size(-1))
        target = sample['target'].view(-1)
        target = model.get_targets(sample, net_output).view(-1)
        loss = F.nll_loss(lprobs, target, size_average=False, ignore_index=self.padding_idx,
                          reduce=reduce)
        sample_size = sample['target'].size(0) if self.args.sentence_avg else sample['ntokens']
@@ -47,6 +49,7 @@ class CrossEntropyCriterion(FairseqCriterion):
        sample_size = sum(log.get('sample_size', 0) for log in logging_outputs)
        agg_output = {
            'loss': loss_sum / sample_size / math.log(2),
            'sample_size': sample_size,
        }
        if sample_size != ntokens:
            agg_output['nll_loss'] = loss_sum / ntokens / math.log(2)
+3 −3
Changes for fairseq/criterions/label_smoothed_cross_entropy.py: 3 added lines, 3 removed lines.
Original line number Diff line number Diff line
@@ -6,8 +6,6 @@
# can be found in the PATENTS file in the same directory.

import math
import torch
import torch.nn.functional as F

from fairseq import utils

@@ -37,7 +35,8 @@ class LabelSmoothedCrossEntropyCriterion(FairseqCriterion):
        """
        net_output = model(**sample['net_input'])
        lprobs = model.get_normalized_probs(net_output, log_probs=True)
        target = sample['target'].unsqueeze(-1)
        lprobs = lprobs.view(-1, lprobs.size(-1))
        target = model.get_targets(sample, net_output).view(-1, 1)
        non_pad_mask = target.ne(self.padding_idx)
        nll_loss = -lprobs.gather(dim=-1, index=target)[non_pad_mask]
        smooth_loss = -lprobs.sum(dim=-1, keepdim=True)[non_pad_mask]
@@ -64,4 +63,5 @@ class LabelSmoothedCrossEntropyCriterion(FairseqCriterion):
        return {
            'loss': sum(log.get('loss', 0) for log in logging_outputs) / sample_size / math.log(2),
            'nll_loss': sum(log.get('nll_loss', 0) for log in logging_outputs) / ntokens / math.log(2),
            'sample_size': sample_size,
        }
+2 −2
Changes for fairseq/data.py: 2 added lines, 2 removed lines.
Original line number Diff line number Diff line
@@ -283,9 +283,9 @@ def _valid_size(src_size, dst_size, max_positions):
        max_src_positions, max_dst_positions = max_positions, max_positions
    else:
        max_src_positions, max_dst_positions = max_positions
    if src_size < 2 or src_size > max_src_positions:
    if src_size < 1 or src_size > max_src_positions:
        return False
    if dst_size is not None and (dst_size < 2 or dst_size > max_dst_positions):
    if dst_size is not None and (dst_size < 1 or dst_size > max_dst_positions):
        return False
    return True

+9 −4
Changes for fairseq/dictionary.py: 9 added lines, 4 removed lines.
Original line number Diff line number Diff line
@@ -6,6 +6,7 @@
# can be found in the PATENTS file in the same directory.

import math
import os
import torch


@@ -23,6 +24,9 @@ class Dictionary(object):
        self.unk_index = self.add_symbol(unk)
        self.nspecial = len(self.symbols)

    def __eq__(self, other):
        return self.indices == other.indices

    def __getitem__(self, idx):
        if idx < len(self.symbols):
            return self.symbols[idx]
@@ -97,8 +101,8 @@ class Dictionary(object):
        """Helper to get index of unk symbol"""
        return self.unk_index

    @staticmethod
    def load(f):
    @classmethod
    def load(cls, f):
        """Loads the dictionary from a text file with the format:

        ```
@@ -111,14 +115,14 @@ class Dictionary(object):
        if isinstance(f, str):
            try:
                with open(f, 'r', encoding='utf-8') as fd:
                    return Dictionary.load(fd)
                    return cls.load(fd)
            except FileNotFoundError as fnfe:
                raise fnfe
            except Exception:
                raise Exception("Incorrect encoding detected in {}, please "
                                "rebuild the dataset".format(f))

        d = Dictionary()
        d = cls()
        for line in f.readlines():
            idx = line.rfind(' ')
            word = line[:idx]
@@ -131,6 +135,7 @@ class Dictionary(object):
    def save(self, f, threshold=3, nwords=-1):
        """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)
        cnt = 0
+17 −8
Changes for fairseq/distributed_utils.py: 17 added lines, 8 removed lines.
Original line number Diff line number Diff line
@@ -10,6 +10,12 @@ import pickle

import torch.distributed

from fairseq import utils


def is_master(args):
    return args.distributed_rank == 0


def distributed_init(args):
    if args.distributed_world_size == 1:
@@ -27,7 +33,7 @@ def distributed_init(args):
            world_size=args.distributed_world_size)

    args.distributed_rank = torch.distributed.get_rank()
    if args.distributed_rank != 0:
    if not is_master(args):
        suppress_output()

    return args.distributed_rank
@@ -104,7 +110,7 @@ def all_gather_list(data, max_size=4096):
    world_size = torch.distributed.get_world_size()
    if not hasattr(all_gather_list, '_in_buffer') or \
            max_size != all_gather_list._in_buffer.size():
        all_gather_list._in_buffer = torch.ByteTensor(max_size)
        all_gather_list._in_buffer = torch.cuda.ByteTensor(max_size)
        all_gather_list._out_buffers = [
            torch.cuda.ByteTensor(max_size)
            for i in range(world_size)
@@ -113,18 +119,21 @@ def all_gather_list(data, max_size=4096):
    out_buffers = all_gather_list._out_buffers

    enc = pickle.dumps(data)
    if len(enc) >= max_size:
        raise ValueError('encoded data exceeds max_size: {}'.format(len(enc)))
    in_buffer[0] = len(enc)
    in_buffer[1:len(enc)+1] = torch.ByteTensor(list(enc))
    enc_size = len(enc)
    if enc_size + 2 > max_size:
        raise ValueError('encoded data exceeds max_size: {}'.format(enc_size + 2))
    assert max_size < 255*256
    in_buffer[0] = enc_size // 255  # this encoding works for max_size < 65k
    in_buffer[1] = enc_size % 255
    in_buffer[2:enc_size+2] = torch.ByteTensor(list(enc))

    torch.distributed.all_gather(out_buffers, in_buffer.cuda())

    result = []
    for i in range(world_size):
        out_buffer = out_buffers[i]
        size = out_buffer[0]
        size = (255 * utils.item(out_buffer[0])) + utils.item(out_buffer[1])
        result.append(
            pickle.loads(bytes(out_buffer[1:size+1].tolist()))
            pickle.loads(bytes(out_buffer[2:size+2].tolist()))
        )
    return result
Loading