构造样本数据 :return: encoder_inputs: [array([0, 0], dtype=int32), array([0, 0], dtype=int32), array([5, 5], dtype=int32), array([7, 7], dtype=int32), array([9, 9], dtype=int32)] decoder_inputs: [array([1, 1], dtype=int32), array([11, 11], dtype=int32), array([
(train_set, batch_num)
| 67 | |
| 68 | |
| 69 | def get_samples(train_set, batch_num): |
| 70 | """构造样本数据 |
| 71 | |
| 72 | :return: |
| 73 | encoder_inputs: [array([0, 0], dtype=int32), array([0, 0], dtype=int32), array([5, 5], dtype=int32), |
| 74 | array([7, 7], dtype=int32), array([9, 9], dtype=int32)] |
| 75 | decoder_inputs: [array([1, 1], dtype=int32), array([11, 11], dtype=int32), array([13, 13], dtype=int32), |
| 76 | array([15, 15], dtype=int32), array([2, 2], dtype=int32)] |
| 77 | """ |
| 78 | # train_set = [[[5, 7, 9], [11, 13, 15, EOS_ID]], [[7, 9, 11], [13, 15, 17, EOS_ID]], [[15, 17, 19], [21, 23, 25, EOS_ID]]] |
| 79 | raw_encoder_input = [] |
| 80 | raw_decoder_input = [] |
| 81 | if batch_num >= len(train_set): |
| 82 | batch_train_set = train_set |
| 83 | else: |
| 84 | random_start = random.randint(0, len(train_set)-batch_num) |
| 85 | batch_train_set = train_set[random_start:random_start+batch_num] |
| 86 | for sample in batch_train_set: |
| 87 | raw_encoder_input.append([PAD_ID] * (input_seq_len - len(sample[0])) + sample[0]) |
| 88 | raw_decoder_input.append([GO_ID] + sample[1] + [PAD_ID] * (output_seq_len - len(sample[1]) - 1)) |
| 89 | |
| 90 | encoder_inputs = [] |
| 91 | decoder_inputs = [] |
| 92 | target_weights = [] |
| 93 | |
| 94 | for length_idx in xrange(input_seq_len): |
| 95 | encoder_inputs.append(np.array([encoder_input[length_idx] for encoder_input in raw_encoder_input], dtype=np.int32)) |
| 96 | for length_idx in xrange(output_seq_len): |
| 97 | decoder_inputs.append(np.array([decoder_input[length_idx] for decoder_input in raw_decoder_input], dtype=np.int32)) |
| 98 | target_weights.append(np.array([ |
| 99 | 0.0 if length_idx == output_seq_len - 1 or decoder_input[length_idx] == PAD_ID else 1.0 for decoder_input in raw_decoder_input |
| 100 | ], dtype=np.float32)) |
| 101 | return encoder_inputs, decoder_inputs, target_weights |
| 102 | |
| 103 | |
| 104 | def seq_to_encoder(input_seq): |