| 20 | from examples.tensorflow.decoder.utils.beam_search import DiverseSiblingSearch |
| 21 | |
| 22 | class 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 | |
| 75 | class DecodingArgument(object): |
| 76 | def __init__( self, |
no outgoing calls
no test coverage detected