MCPcopy Create free account
hub / github.com/NVIDIA/FasterTransformer / main

Function main

examples/pytorch/decoding/decoding_example.py:32–191  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

30from examples.pytorch.decoding.utils.ft_decoding import FtDecodingWeights, CustomDecoding
31
32def main():
33 parser = argparse.ArgumentParser()
34 parser.add_argument('batch_size', type=int,
35 help='batch size')
36 parser.add_argument('layer_num', type=int,
37 help='number of layers')
38 parser.add_argument('seq_len', type=int,
39 help='sequence length')
40 parser.add_argument('head_num', type=int,
41 help='head number')
42 parser.add_argument('head_size', type=int,
43 help='size per head')
44 parser.add_argument('-inter_size', '--inter_size', type=int, default=0, metavar='NUMBER',
45 help='inter_size (default: 0)')
46 parser.add_argument('-mem_hidden', '--memory_hidden_dim', type=int, default=512, metavar='NUMBER',
47 help='memory hidden dim (default: 512)')
48 parser.add_argument('beam_size', type=int,
49 help='beam size')
50 parser.add_argument('vocab_size', type=int,
51 help='vocab size')
52 parser.add_argument('--data_type', type=str, choices=['fp32', 'fp16', 'bf16'], default='fp32')
53 parser.add_argument('--time', action='store_true',
54 help='test the time or not.')
55 parser.add_argument('--use_pretrained', action='store_true',
56 help='use pretrained weights or not.')
57 parser.add_argument('--decoding_ths_path', type=str, default='./lib/libth_transformer.so',
58 help='path of the pyt_fastertransformer dynamic lib file')
59 parser.add_argument('--decoder_ths_path', type=str, default='./lib/libth_transformer.so',
60 help='path of the pyt_fastertransformer dynamic lib file')
61 parser.add_argument('-diversity_rate', '--beam_search_diversity_rate', type=float, default=0.0, metavar='NUMBER',
62 help='deviersity rate of beam search. default is 0. When diversity rate = 0, it is equivalent to the naive beam search.')
63 parser.add_argument('-topk', '--sampling_topk', type=int, default=1, metavar='NUMBER',
64 help='Candidate (k) value of top k sampling in decoding. Default is 1.')
65 parser.add_argument('-topp', '--sampling_topp', type=float, default=0.0, metavar='NUMBER',
66 help='Probability (p) value of top p sampling in decoding. Default is 0.0. ')
67
68 args = parser.parse_args()
69
70 torch.manual_seed(0)
71 random.seed(0)
72 np.random.seed(0)
73
74 if args.use_pretrained:
75 layer_num = 6
76 head_num = 8
77 head_size = 64
78 inter_size = head_num * head_size * 4
79 vocab_size = 31538
80 else:
81 layer_num = args.layer_num
82 head_num = args.head_num
83 head_size = args.head_size
84 inter_size = args.inter_size
85 if inter_size == 0:
86 inter_size = 4 * head_num * head_size
87 vocab_size = args.vocab_size
88 hidden_dim = head_num * head_size
89 start_id = 2

Callers 1

Calls 14

to_cudaMethod · 0.95
to_halfMethod · 0.95
to_bfloat16Method · 0.95
ArgHelperClass · 0.90
DecodingWeightsClass · 0.90
FtDecodingWeightsClass · 0.90
TorchDecodingClass · 0.90
CustomDecodingClass · 0.90
maxMethod · 0.80
minMethod · 0.80
fix_keyFunction · 0.70
cudaMethod · 0.45

Tested by

no test coverage detected