| 166 | |
| 167 | # wrapper function |
| 168 | def ftt5_encoder(inputs, seq_len, encoder_params): |
| 169 | transformer_op_module = tf.load_op_library(os.path.join('./lib/libtf_t5.so')) |
| 170 | |
| 171 | outputs = transformer_op_module.t5_encoder(inputs, |
| 172 | seq_len, |
| 173 | *encoder_params.weights, |
| 174 | head_num = encoder_params.num_heads, |
| 175 | head_size = encoder_params.head_size, # encoder_config.d_kv |
| 176 | inter_size = encoder_params.inter_size, # encoder_config.d_ff, |
| 177 | num_layer = encoder_params.num_layer, |
| 178 | d_model = encoder_params.d_model, |
| 179 | num_bucket = encoder_params.num_bucket, |
| 180 | max_distance = encoder_params.max_distance, |
| 181 | remove_padding = True, |
| 182 | t5_with_bias = encoder_params.t5_with_bias, |
| 183 | activation_type = encoder_params.activation_type, |
| 184 | q_scaling = encoder_params.q_scaling, |
| 185 | position_embedding_type=encoder_params.position_embedding_type) |
| 186 | |
| 187 | return outputs |