You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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:

  1. 安装TFP:
pip install tensorflow-probability
  1. 修改采样代码:
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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.21 10:39:16