Commit 30ef667d authored by Angela Fan's avatar Angela Fan
Browse files

add model override argument from load_ensemble_for_inference at generation...

add model override argument from load_ensemble_for_inference at generation time, updating readme for stories
parent ff3db3cd
Loading
Loading
Loading
Loading
+3 −1
Changes for examples/stories/README.md: 3 added lines, 1 removed line.
Original line number Diff line number Diff line
@@ -26,5 +26,7 @@ $ python train.py data-bin/writingPrompts -a fconv_self_att_wp --lr 0.25 --clip-
# add the arguments: --pretrained True --pretrained-checkpoint path/to/checkpoint

# Generate:
$ python generate.py data-bin/writingPrompts --path /path/to/trained/model/checkpoint_best.pt --batch-size 32 --beam 1 --sampling --sampling-topk 10 --sampling-temperature 0.8 --nbest 1 
# Note: to load the pretrained model at generation time, you need to pass in a model-override argument to communicate to the fusion model at generation time where you have placed the pretrained checkpoint. By default, it will load the exact path of the fusion model's pretrained model from training time. You should use model-override if you have moved the pretrained model (or are using our provided models). If you are generating from a non-fusion model, the model-override argument is not necessary.

$ python generate.py data-bin/writingPrompts --path /path/to/trained/model/checkpoint_best.pt --batch-size 32 --beam 1 --sampling --sampling-topk 10 --sampling-temperature 0.8 --nbest 1 --model-overrides "{'pretrained_checkpoint':'/path/to/pretrained/model/checkpoint'}"
```
+1 −1
Changes for fairseq/models/fconv_self_att.py: 1 added line, 1 removed line.
Original line number Diff line number Diff line
@@ -81,7 +81,7 @@ class FConvModelSelfAtt(FairseqModel):
        trained_encoder, trained_decoder = None, None
        pretrained = eval(args.pretrained)
        if pretrained:
            print("| Loading pretrained model")
            print("| loading pretrained model")
            trained_model = utils.load_ensemble_for_inference(
                # not actually for inference, but loads pretrained model parameters
                filenames=[args.pretrained_checkpoint],
+2 −0
Changes for fairseq/options.py: 2 added lines, 0 removed lines.
Original line number Diff line number Diff line
@@ -290,6 +290,8 @@ def add_generation_args(parser):
                       help='sample from top K likely next words instead of all words')
    group.add_argument('--sampling-temperature', default=1, type=float, metavar='N',
                       help='temperature for random sampling')
    group.add_argument('--model-overrides', default="{}", type=str, metavar='DICT',
                       help='a dictionary used to override model args at generation that were used during model training')
    return group


+1 −1
Changes for generate.py: 1 added line, 1 removed line.
Original line number Diff line number Diff line
@@ -38,7 +38,7 @@ def main(args):

    # Load ensemble
    print('| loading model(s) from {}'.format(args.path))
    models, _ = utils.load_ensemble_for_inference(args.path.split(':'), task)
    models, _ = utils.load_ensemble_for_inference(args.path.split(':'), task, model_arg_overrides=eval(args.model_overrides))

    # Optimize ensemble for generation
    for model in models: