MCPcopy Create free account
hub / github.com/XL2248/MSCTD / _sampling_step

Function _sampling_step

src_code/thumt1_code/thumt/utils/sampling.py:53–101  ·  view source on GitHub ↗
(time, func, state, min_length, max_length, pad_id, eos_id)

Source from the content-addressed store, hash-verified

51
52
53def _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
104def random_sample(func, state, batch_size, min_length, max_length, pad_id,

Callers 1

_loop_fnFunction · 0.70

Calls 1

SamplerStateClass · 0.70

Tested by

no test coverage detected