TensorFlow v1.x教程报错:__init__()收到意外参数'num_samples'
TensorFlow 2.x 解决文本生成中Multinomial采样报错的方案
报错原因
你调用tf.compat.v1.distributions.Multinomial时传入了num_samples参数,但该类的构造函数并不接受这个参数。原代码想实现从预测概率分布中采样1个字符ID的逻辑,TF1.x的旧写法在TF2.x中不再适用。
TF2.x 解决方案
方案1:使用tf.random.categorical(推荐,无需额外依赖)
直接用TF2.x原生的tf.random.categorical函数实现多项分布采样,替代原来的Multinomial调用:
修改generate_text函数中的采样代码部分:
# 替换原来的predicted_id行 predictions = predictions / temperature # tf.random.categorical接受2D张量,因此将1D的predictions扩展维度 predicted_id = tf.random.categorical(tf.expand_dims(predictions, 0), num_samples=1)[0, 0].numpy()
完整修改后的generate_text函数:
def generate_text(model, start_string): # Evaluation step (generating text using the learned model) # Number of characters to generate num_generate = 1000 # Converting our start string to numbers (vectorizing) input_eval = [char2idx[s] for s in start_string] input_eval = tf.expand_dims(input_eval, 0) # Empty string to store our results text_generated = [] # Low temperatures results in more predictable text. # Higher temperatures results in more surprising text. # Experiment to find the best setting. temperature = 1.0 # Here batch size == 1 model.reset_states() for i in range(num_generate): predictions = model(input_eval) # remove the batch dimension predictions = tf.squeeze(predictions, 0) # using a multinomial distribution to predict the word returned by the model predictions = predictions / temperature # 使用TF2.x原生的categorical采样 predicted_id = tf.random.categorical(tf.expand_dims(predictions, 0), num_samples=1)[0, 0].numpy() # We pass the predicted word as the next input to the model # along with the previous hidden state input_eval = tf.expand_dims([predicted_id], 0) text_generated.append(idx2char[predicted_id]) return (start_string + ''.join(text_generated))
方案2:使用TensorFlow Probability的Multinomial(需安装TFP)
若倾向于使用分布类写法,可安装TensorFlow Probability库后用tfp.distributions.Multinomial:
- 安装TFP:
pip install tensorflow-probability
- 修改采样代码:
import tensorflow_probability as tfp # 替换原来的predicted_id行 predictions = predictions / temperature # 创建Multinomial分布,total_count=1表示采样1次 multinomial = tfp.distributions.Multinomial(total_count=1, logits=predictions) # 采样并提取结果 predicted_id = tf.argmax(multinomial.sample(1), axis=-1).numpy()[0]
关键说明
tf.random.categorical是TF2.x专门用于分类分布采样的原生函数,输入为未归一化的logits,num_samples指定采样数量,用法更简洁直接。- 升级到TF2.x后,建议尽量避免使用
tf.compat.v1下的旧API,改用原生TF2.x接口,减少兼容层带来的潜在问题。
内容的提问来源于stack exchange,提问作者Black Snow
相关产品推荐
相关产品推荐

