Loading fairseq/sequence_generator.py +9 −9 Original line number Diff line number Diff line Loading @@ -83,10 +83,11 @@ class SequenceGenerator(object): timer.start() with torch.no_grad(): hypos = self.generate( input['src_tokens'], input['src_lengths'], beam_size=beam_size, maxlen=int(maxlen_a*srclen + maxlen_b), prefix_tokens=s['target'][:, :prefix_size] if prefix_size > 0 else None, **net_input, ) if timer is not None: timer.stop(sum(len(h[0]['tokens']) for h in hypos)) Loading @@ -96,13 +97,12 @@ class SequenceGenerator(object): ref = utils.strip_pad(s['target'].data[i, :], self.pad) if s['target'] is not None else None yield id, src, ref, hypos[i] def generate(self, beam_size=None, maxlen=None, prefix_tokens=None, **net_input): def generate(self, src_tokens, src_lengths, beam_size=None, maxlen=None, prefix_tokens=None): """Generate a batch of translations.""" with torch.no_grad(): return self._generate(beam_size, maxlen, prefix_tokens, **net_input) return self._generate(src_tokens, src_lengths, beam_size, maxlen, prefix_tokens) def _generate(self, beam_size=None, maxlen=None, prefix_tokens=None, **net_input): src_tokens = net_input['src_tokens'] def _generate(self, src_tokens, src_lengths, beam_size=None, maxlen=None, prefix_tokens=None): bsz, srclen = src_tokens.size() maxlen = min(maxlen, self.maxlen) if maxlen is not None else self.maxlen Loading @@ -121,10 +121,10 @@ class SequenceGenerator(object): incremental_states[model] = None # compute the encoder output for each beam encoder_out = model.encoder(**net_input) new_order = torch.arange(bsz).view(-1, 1).repeat(1, beam_size).view(-1) new_order = new_order.to(net_input['src_tokens'].device) encoder_out = model.encoder.reorder_encoder_out(encoder_out, new_order) encoder_out = model.encoder( src_tokens.repeat(1, beam_size).view(-1, srclen), src_lengths.expand(beam_size, src_lengths.numel()).t().contiguous().view(-1), ) encoder_outs.append(encoder_out) # initialize buffers Loading tests/test_sequence_generator.py +6 −6 Original line number Diff line number Diff line Loading @@ -85,7 +85,7 @@ class TestSequenceGenerator(unittest.TestCase): def test_with_normalization(self): generator = SequenceGenerator([self.model], self.tgt_dict) hypos = generator.generate(src_tokens=self.src_tokens, src_lengths=self.src_lengths, beam_size=2) hypos = generator.generate(self.src_tokens, self.src_lengths, beam_size=2) eos, w1, w2 = self.eos, self.w1, self.w2 # sentence 1, beam 1 self.assertHypoTokens(hypos[0][0], [w1, eos]) Loading @@ -104,7 +104,7 @@ class TestSequenceGenerator(unittest.TestCase): # Sentence 1: unchanged from the normalized case # Sentence 2: beams swap order generator = SequenceGenerator([self.model], self.tgt_dict, normalize_scores=False) hypos = generator.generate(src_tokens=self.src_tokens, src_lengths=self.src_lengths, beam_size=2) hypos = generator.generate(self.src_tokens, self.src_lengths, beam_size=2) eos, w1, w2 = self.eos, self.w1, self.w2 # sentence 1, beam 1 self.assertHypoTokens(hypos[0][0], [w1, eos]) Loading @@ -122,7 +122,7 @@ class TestSequenceGenerator(unittest.TestCase): def test_with_lenpen_favoring_short_hypos(self): lenpen = 0.6 generator = SequenceGenerator([self.model], self.tgt_dict, len_penalty=lenpen) hypos = generator.generate(src_tokens=self.src_tokens, src_lengths=self.src_lengths, beam_size=2) hypos = generator.generate(self.src_tokens, self.src_lengths, beam_size=2) eos, w1, w2 = self.eos, self.w1, self.w2 # sentence 1, beam 1 self.assertHypoTokens(hypos[0][0], [w1, eos]) Loading @@ -140,7 +140,7 @@ class TestSequenceGenerator(unittest.TestCase): def test_with_lenpen_favoring_long_hypos(self): lenpen = 5.0 generator = SequenceGenerator([self.model], self.tgt_dict, len_penalty=lenpen) hypos = generator.generate(src_tokens=self.src_tokens, src_lengths=self.src_lengths, beam_size=2) hypos = generator.generate(self.src_tokens, self.src_lengths, beam_size=2) eos, w1, w2 = self.eos, self.w1, self.w2 # sentence 1, beam 1 self.assertHypoTokens(hypos[0][0], [w2, w1, w2, eos]) Loading @@ -157,7 +157,7 @@ class TestSequenceGenerator(unittest.TestCase): def test_maxlen(self): generator = SequenceGenerator([self.model], self.tgt_dict, maxlen=2) hypos = generator.generate(src_tokens=self.src_tokens, src_lengths=self.src_lengths, beam_size=2) hypos = generator.generate(self.src_tokens, self.src_lengths, beam_size=2) eos, w1, w2 = self.eos, self.w1, self.w2 # sentence 1, beam 1 self.assertHypoTokens(hypos[0][0], [w1, eos]) Loading @@ -174,7 +174,7 @@ class TestSequenceGenerator(unittest.TestCase): def test_no_stop_early(self): generator = SequenceGenerator([self.model], self.tgt_dict, stop_early=False) hypos = generator.generate(src_tokens=self.src_tokens, src_lengths=self.src_lengths, beam_size=2) hypos = generator.generate(self.src_tokens, self.src_lengths, beam_size=2) eos, w1, w2 = self.eos, self.w1, self.w2 # sentence 1, beam 1 self.assertHypoTokens(hypos[0][0], [w1, eos]) Loading Loading
fairseq/sequence_generator.py +9 −9 Original line number Diff line number Diff line Loading @@ -83,10 +83,11 @@ class SequenceGenerator(object): timer.start() with torch.no_grad(): hypos = self.generate( input['src_tokens'], input['src_lengths'], beam_size=beam_size, maxlen=int(maxlen_a*srclen + maxlen_b), prefix_tokens=s['target'][:, :prefix_size] if prefix_size > 0 else None, **net_input, ) if timer is not None: timer.stop(sum(len(h[0]['tokens']) for h in hypos)) Loading @@ -96,13 +97,12 @@ class SequenceGenerator(object): ref = utils.strip_pad(s['target'].data[i, :], self.pad) if s['target'] is not None else None yield id, src, ref, hypos[i] def generate(self, beam_size=None, maxlen=None, prefix_tokens=None, **net_input): def generate(self, src_tokens, src_lengths, beam_size=None, maxlen=None, prefix_tokens=None): """Generate a batch of translations.""" with torch.no_grad(): return self._generate(beam_size, maxlen, prefix_tokens, **net_input) return self._generate(src_tokens, src_lengths, beam_size, maxlen, prefix_tokens) def _generate(self, beam_size=None, maxlen=None, prefix_tokens=None, **net_input): src_tokens = net_input['src_tokens'] def _generate(self, src_tokens, src_lengths, beam_size=None, maxlen=None, prefix_tokens=None): bsz, srclen = src_tokens.size() maxlen = min(maxlen, self.maxlen) if maxlen is not None else self.maxlen Loading @@ -121,10 +121,10 @@ class SequenceGenerator(object): incremental_states[model] = None # compute the encoder output for each beam encoder_out = model.encoder(**net_input) new_order = torch.arange(bsz).view(-1, 1).repeat(1, beam_size).view(-1) new_order = new_order.to(net_input['src_tokens'].device) encoder_out = model.encoder.reorder_encoder_out(encoder_out, new_order) encoder_out = model.encoder( src_tokens.repeat(1, beam_size).view(-1, srclen), src_lengths.expand(beam_size, src_lengths.numel()).t().contiguous().view(-1), ) encoder_outs.append(encoder_out) # initialize buffers Loading
tests/test_sequence_generator.py +6 −6 Original line number Diff line number Diff line Loading @@ -85,7 +85,7 @@ class TestSequenceGenerator(unittest.TestCase): def test_with_normalization(self): generator = SequenceGenerator([self.model], self.tgt_dict) hypos = generator.generate(src_tokens=self.src_tokens, src_lengths=self.src_lengths, beam_size=2) hypos = generator.generate(self.src_tokens, self.src_lengths, beam_size=2) eos, w1, w2 = self.eos, self.w1, self.w2 # sentence 1, beam 1 self.assertHypoTokens(hypos[0][0], [w1, eos]) Loading @@ -104,7 +104,7 @@ class TestSequenceGenerator(unittest.TestCase): # Sentence 1: unchanged from the normalized case # Sentence 2: beams swap order generator = SequenceGenerator([self.model], self.tgt_dict, normalize_scores=False) hypos = generator.generate(src_tokens=self.src_tokens, src_lengths=self.src_lengths, beam_size=2) hypos = generator.generate(self.src_tokens, self.src_lengths, beam_size=2) eos, w1, w2 = self.eos, self.w1, self.w2 # sentence 1, beam 1 self.assertHypoTokens(hypos[0][0], [w1, eos]) Loading @@ -122,7 +122,7 @@ class TestSequenceGenerator(unittest.TestCase): def test_with_lenpen_favoring_short_hypos(self): lenpen = 0.6 generator = SequenceGenerator([self.model], self.tgt_dict, len_penalty=lenpen) hypos = generator.generate(src_tokens=self.src_tokens, src_lengths=self.src_lengths, beam_size=2) hypos = generator.generate(self.src_tokens, self.src_lengths, beam_size=2) eos, w1, w2 = self.eos, self.w1, self.w2 # sentence 1, beam 1 self.assertHypoTokens(hypos[0][0], [w1, eos]) Loading @@ -140,7 +140,7 @@ class TestSequenceGenerator(unittest.TestCase): def test_with_lenpen_favoring_long_hypos(self): lenpen = 5.0 generator = SequenceGenerator([self.model], self.tgt_dict, len_penalty=lenpen) hypos = generator.generate(src_tokens=self.src_tokens, src_lengths=self.src_lengths, beam_size=2) hypos = generator.generate(self.src_tokens, self.src_lengths, beam_size=2) eos, w1, w2 = self.eos, self.w1, self.w2 # sentence 1, beam 1 self.assertHypoTokens(hypos[0][0], [w2, w1, w2, eos]) Loading @@ -157,7 +157,7 @@ class TestSequenceGenerator(unittest.TestCase): def test_maxlen(self): generator = SequenceGenerator([self.model], self.tgt_dict, maxlen=2) hypos = generator.generate(src_tokens=self.src_tokens, src_lengths=self.src_lengths, beam_size=2) hypos = generator.generate(self.src_tokens, self.src_lengths, beam_size=2) eos, w1, w2 = self.eos, self.w1, self.w2 # sentence 1, beam 1 self.assertHypoTokens(hypos[0][0], [w1, eos]) Loading @@ -174,7 +174,7 @@ class TestSequenceGenerator(unittest.TestCase): def test_no_stop_early(self): generator = SequenceGenerator([self.model], self.tgt_dict, stop_early=False) hypos = generator.generate(src_tokens=self.src_tokens, src_lengths=self.src_lengths, beam_size=2) hypos = generator.generate(self.src_tokens, self.src_lengths, beam_size=2) eos, w1, w2 = self.eos, self.w1, self.w2 # sentence 1, beam 1 self.assertHypoTokens(hypos[0][0], [w1, eos]) Loading