从输入空格分隔的数字id串,转成预测用的encoder、decoder、target_weight等
(input_seq)
| 102 | |
| 103 | |
| 104 | def seq_to_encoder(input_seq): |
| 105 | """从输入空格分隔的数字id串,转成预测用的encoder、decoder、target_weight等 |
| 106 | """ |
| 107 | input_seq_array = [int(v) for v in input_seq.split()] |
| 108 | encoder_input = [PAD_ID] * (input_seq_len - len(input_seq_array)) + input_seq_array |
| 109 | decoder_input = [GO_ID] + [PAD_ID] * (output_seq_len - 1) |
| 110 | encoder_inputs = [np.array([v], dtype=np.int32) for v in encoder_input] |
| 111 | decoder_inputs = [np.array([v], dtype=np.int32) for v in decoder_input] |
| 112 | target_weights = [np.array([1.0], dtype=np.float32)] * output_seq_len |
| 113 | return encoder_inputs, decoder_inputs, target_weights |
| 114 | |
| 115 | |
| 116 | def get_model(feed_previous=False): |