()
| 30 | from examples.pytorch.decoding.utils.decoding import DecodingWeights |
| 31 | |
| 32 | def 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 |
no test coverage detected