Generates translations of a given source sentence. Args: models (List[~fairseq.models.FairseqModel]): ensemble of models, currently support fairseq.models.TransformerModel for scripting beam_size (int, optional): beam width (default: 1) ma
(
self,
models,
tgt_dict,
beam_size=1,
max_len_a=0,
max_len_b=200,
max_len=0,
min_len=1,
normalize_scores=True,
len_penalty=1.0,
unk_penalty=0.0,
temperature=1.0,
match_source_len=False,
no_repeat_ngram_size=0,
search_strategy=None,
eos=None,
symbols_to_strip_from_output=None,
lm_model=None,
lm_weight=1.0,
constraint_trie=None,
constraint_range=None,
gen_code=False,
gen_box=False,
ignore_eos=False,
zero_shot=False
)
| 18 | |
| 19 | class SequenceGenerator(nn.Module): |
| 20 | def __init__( |
| 21 | self, |
| 22 | models, |
| 23 | tgt_dict, |
| 24 | beam_size=1, |
| 25 | max_len_a=0, |
| 26 | max_len_b=200, |
| 27 | max_len=0, |
| 28 | min_len=1, |
| 29 | normalize_scores=True, |
| 30 | len_penalty=1.0, |
| 31 | unk_penalty=0.0, |
| 32 | temperature=1.0, |
| 33 | match_source_len=False, |
| 34 | no_repeat_ngram_size=0, |
| 35 | search_strategy=None, |
| 36 | eos=None, |
| 37 | symbols_to_strip_from_output=None, |
| 38 | lm_model=None, |
| 39 | lm_weight=1.0, |
| 40 | constraint_trie=None, |
| 41 | constraint_range=None, |
| 42 | gen_code=False, |
| 43 | gen_box=False, |
| 44 | ignore_eos=False, |
| 45 | zero_shot=False |
| 46 | ): |
| 47 | """Generates translations of a given source sentence. |
| 48 | |
| 49 | Args: |
| 50 | models (List[~fairseq.models.FairseqModel]): ensemble of models, |
| 51 | currently support fairseq.models.TransformerModel for scripting |
| 52 | beam_size (int, optional): beam width (default: 1) |
| 53 | max_len_a/b (int, optional): generate sequences of maximum length |
| 54 | ax + b, where x is the source length |
| 55 | max_len (int, optional): the maximum length of the generated output |
| 56 | (not including end-of-sentence) |
| 57 | min_len (int, optional): the minimum length of the generated output |
| 58 | (not including end-of-sentence) |
| 59 | normalize_scores (bool, optional): normalize scores by the length |
| 60 | of the output (default: True) |
| 61 | len_penalty (float, optional): length penalty, where <1.0 favors |
| 62 | shorter, >1.0 favors longer sentences (default: 1.0) |
| 63 | unk_penalty (float, optional): unknown word penalty, where <0 |
| 64 | produces more unks, >0 produces fewer (default: 0.0) |
| 65 | temperature (float, optional): temperature, where values |
| 66 | >1.0 produce more uniform samples and values <1.0 produce |
| 67 | sharper samples (default: 1.0) |
| 68 | match_source_len (bool, optional): outputs should match the source |
| 69 | length (default: False) |
| 70 | """ |
| 71 | super().__init__() |
| 72 | if isinstance(models, EnsembleModel): |
| 73 | self.model = models |
| 74 | else: |
| 75 | self.model = EnsembleModel(models) |
| 76 | self.gen_code = gen_code |
| 77 | self.gen_box = gen_box |
nothing calls this directly
no test coverage detected