()
| 30 | from examples.pytorch.decoding.utils.ft_decoding import FtDecodingWeights, CustomDecoding |
| 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('-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 |
no test coverage detected