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

Function main

examples/pytorch/decoder/decoder_example.py:32–159  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

30from examples.pytorch.decoding.utils.decoding import DecodingWeights
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('--step', type=int, default=0,
45 help='decoding step number')
46 parser.add_argument('--decoder_ths_path', type=str, default='./lib/libth_transformer.so',
47 help='path of the pyt_fastertransformer dynamic lib file')
48 parser.add_argument('--time', action='store_true',
49 help='test the time or not.')
50 parser.add_argument('--ths_path', type=str, default='./lib/libth_transformer.so',
51 help='path of the pyt_fastertransformer dynamic lib file')
52 parser.add_argument('-d', '--data_type', type=str, default="fp32", metavar='STRING',
53 help='data type (default: fp32)', choices=['fp32', 'fp16', 'bf16'])
54
55 args = parser.parse_args()
56
57 hidden_dim = args.head_num * args.head_size
58
59 if args.step <= 0:
60 step = args.seq_len
61 else:
62 step = args.step
63
64 print("\n=============== Argument ===============")
65 print('batch_size: ' + str(args.batch_size))
66 print('layer_num: ' + str(args.layer_num))
67 print('seq_len: ' + str(args.seq_len))
68 print('head_num: ' + str(args.head_num))
69 print('head_size: ' + str(args.head_size))
70 print('hidden_dim: ' + str(hidden_dim))
71 print('step: ' + str(step))
72 print('data_type: ' + str(args.data_type))
73 print('test_time: ' + str(args.time))
74 print("========================================\n")
75
76 np.random.seed(1)
77 torch.manual_seed(0)
78 random.seed(0)
79
80 inp = torch.empty(args.batch_size, 1, hidden_dim).cuda()
81 mem = torch.empty(args.batch_size, args.seq_len, hidden_dim).cuda() # We assume mem_hidden_dim = hidden_dim
82 torch.nn.init.uniform_(inp, -0.5, 0.5)
83 torch.nn.init.uniform_(mem, -0.5, 0.5)
84 if args.data_type == 'fp16':
85 inp = inp.half()
86 mem = mem.half()
87 mem_seq_lens = torch.randint(1, args.seq_len+1, (args.batch_size,), dtype=torch.int32).cuda()
88 src_pad_mask = ~sequence_mask(mem_seq_lens, args.seq_len).unsqueeze(1)
89

Callers 1

decoder_example.pyFile · 0.70

Calls 15

to_cudaMethod · 0.95
to_cudaMethod · 0.95
to_halfMethod · 0.95
to_halfMethod · 0.95
to_bfloat16Method · 0.95
to_bfloat16Method · 0.95
DecodingWeightsClass · 0.90
FtDecoderWeightsClass · 0.90
ONMTDecoderClass · 0.90
FTDecoderClass · 0.90
init_op_cacheFunction · 0.90
init_onmt_cacheFunction · 0.90

Tested by

no test coverage detected