| 51 | |
| 52 | |
| 53 | def _sampling_step(time, func, state, min_length, max_length, pad_id, eos_id): |
| 54 | # Compute log probabilities |
| 55 | seqs = state.inputs |
| 56 | # [batch_size * num_samples, vocab_size] |
| 57 | step_log_probs, next_state = func(seqs, state.state) |
| 58 | |
| 59 | # Suppress <eos> if needed |
| 60 | batch_size = tf.shape(step_log_probs)[0] |
| 61 | vocab_size = step_log_probs.shape[-1].value or tf.shape(step_log_probs)[1] |
| 62 | add_mask = tf.one_hot(eos_id, vocab_size, dtype=step_log_probs.dtype, |
| 63 | on_value=step_log_probs.dtype.min, |
| 64 | off_value=0.0) |
| 65 | add_mask = utils.tile_batch(tf.reshape(add_mask, [1, -1]), batch_size) |
| 66 | add_mask = tf.where(time < min_length, add_mask, tf.zeros_like(add_mask)) |
| 67 | step_log_probs = step_log_probs + add_mask |
| 68 | |
| 69 | # sample from distribution |
| 70 | symbol_indices = tf.multinomial(step_log_probs, 1, output_dtype=tf.int32) |
| 71 | symbol_scores = tf.squeeze(utils.gather_2d(step_log_probs, symbol_indices), |
| 72 | axis=1) |
| 73 | curr_flags = tf.squeeze(tf.equal(symbol_indices, eos_id), axis=1) |
| 74 | curr_flags = tf.logical_or(state.flags, curr_flags) |
| 75 | |
| 76 | # Append <pad> to finished samples |
| 77 | symbol_indices = tf.where(state.flags, tf.fill([batch_size, 1], pad_id), |
| 78 | symbol_indices) |
| 79 | symbol_scores = tf.where(state.flags, tf.zeros([batch_size]), |
| 80 | symbol_scores) |
| 81 | |
| 82 | # Force sampler to generate <eos> if length exceed max_length |
| 83 | eos_flags = tf.where(time > max_length, tf.ones([batch_size], tf.bool), |
| 84 | tf.zeros([batch_size], tf.bool)) |
| 85 | eos_scores = tf.squeeze(utils.gather_2d(step_log_probs, |
| 86 | tf.fill([batch_size, 1], eos_id)), |
| 87 | axis=1) |
| 88 | eos_indices = tf.fill([batch_size, 1], eos_id) |
| 89 | cond = tf.logical_and(tf.logical_not(curr_flags), eos_flags) |
| 90 | curr_flags = tf.logical_or(curr_flags, eos_flags) |
| 91 | symbol_indices = tf.where(cond, eos_indices, symbol_indices) |
| 92 | symbol_scores = tf.where(cond, eos_scores, symbol_scores) |
| 93 | |
| 94 | new_state = SamplerState( |
| 95 | inputs=tf.concat([seqs, symbol_indices], axis=1), |
| 96 | state=next_state, |
| 97 | scores=state.scores + symbol_scores, |
| 98 | flags=curr_flags |
| 99 | ) |
| 100 | |
| 101 | return time + 1, new_state |
| 102 | |
| 103 | |
| 104 | def random_sample(func, state, batch_size, min_length, max_length, pad_id, |