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

Class TransformerArgument

examples/tensorflow/decoder/utils/common.py:22–73  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

20from examples.tensorflow.decoder.utils.beam_search import DiverseSiblingSearch
21
22class TransformerArgument:
23 def __init__( self,
24 beam_width,
25 head_num,
26 size_per_head,
27 inter_size,
28 num_layer,
29 dtype=tf.float32,
30 kernel_init_range=0.02,
31 bias_init_range=0.02,
32 fuse_qkv=True,
33 remove_padding=False,
34 int8_mode=0,
35 allow_gemm_test=False,
36 memory_hidden_dim=-1):
37 '''
38 The arguments of Transformer layer (for both encoder and decoder).
39
40 Args:
41 beam_width: The beam_width size for beam search. This argument is always one for encoder.
42 head_num: The head number of self attention in transformer layer.
43 size_per_head: The size of hidden dimension for each head of self attention in transformer layer.
44 inter_size: The size of intermediate dimension of FFN layer.
45 num_layer: The number of transformer layer. For example, BERT-base uses 12 layers.
46 dtype: The data type of weights initializer and inputs.
47 kernel_init_range: The initializer range of kernel for all convolution layer and fully-connected layer.
48 kernel_init_range: The initializer range of bias for all convolution layer and fully-connected layer.
49 fuse_qkv: bool. Whether fuse the q, k, v gemm or not.
50 remove_padding: bool. Remove the padding of sentences of encoder.
51 int8_mode: Mode of int8 quantization. 0 means not using int8 quantization, 1 means using int8 quantization without quantizing residuals,
52 2 means using int8 quantization with quantizing residuals.
53 allow_gemm_test: whether allow gemm test inside FT.
54 '''
55
56 self.beam_width = beam_width
57 self.head_num = head_num
58 self.size_per_head = size_per_head
59 self.inter_size = inter_size
60 self.num_layer = num_layer
61 self.dtype = dtype
62 self.hidden_dim = self.head_num * self.size_per_head
63 self.kernel_init_range = kernel_init_range
64 self.bias_init_range = bias_init_range
65 self.int8_mode = int8_mode
66 self.allow_gemm_test = allow_gemm_test
67 if self.dtype == tf.float32:
68 self.check_threshold = 2e-5
69 elif self.dtype == tf.float16:
70 self.check_threshold = 2e-2
71 self.fuse_qkv = fuse_qkv
72 self.remove_padding = remove_padding
73 self.memory_hidden_dim = memory_hidden_dim
74
75class DecodingArgument(object):
76 def __init__( self,

Callers 3

translateFunction · 0.90
decoder_example.pyFile · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected