【问题标题】:How to output the list of probabilities on each token via model.generate?如何通过model.generate输出每个token的概率列表?
【发布时间】:2023-01-19 12:44:22
【问题描述】:

现在我有:

model = GPTNeoForCausalLM.from_pretrained(model_name)
tokenizer = GPT2Tokenizer.from_pretrained(model_name)
input_ids = tokenizer(prompt, return_tensors="pt").input_ids.cuda()
gen_tokens = model.generate(input_ids, do_sample=specifiedDoSample, output_scores=True, temperature=specifiedTemperature, max_new_tokens=specifiedNumTokens, repetition_penalty=specifiedRepetitionPenalty, top_p=specifiedTopP)
gen_text = tokenizer.batch_decode(gen_tokens)[0]
print(gen_text)

这将打印生成的文本。但是,我希望它列出每个步骤中的前 N ​​个标记及其概率(N 是我指定的数字),类似于 OpenAI 的 beta 游乐场,您可以在其中选择“显示概率:全谱”。例如,如果提示是“你现在是一个”,下一个标记应该是类似{“vampire”: 51%, “corpse”: 32% ...等}

通过 Huggingface 变形金刚做到这一点的最简单方法是什么?

【问题讨论】:

    标签: python nlp huggingface-transformers gpt-3


    【解决方案1】:

    您需要在对生成方法的调用中添加“, output_scores=True, return_dict_in_generate=True”,这将为您提供生成短语的每个字符的分数表,其中包含一个带有分数的张量(需要 softmax 以获得概率) 在波束搜索中每个可能序列的每个标记。

    查看 transformers 源代码树中的 generation_utils.py,从“def generate”开始

    【讨论】:

    • 正如目前所写,您的答案尚不清楚。请edit 添加更多详细信息,以帮助其他人了解这如何解决所提出的问题。你可以找到更多关于如何写出好的答案的信息in the help center。
    • 谢谢。我不需要指定集束搜索或采样以及运行次数吗?比如说,获得前 50 个下一个代币。我遇到了这个问题:github.com/huggingface/transformers/issues/10012 我可以使用束搜索来获得最佳选择,但概率是错误的
    • 光束采样参数在模型中是默认的。您可以添加 num_beams、num_beam_groups(不知道这是做什么的)、num_return_sequence 作为运行次数。还有很多其他参数,例如 n_gram interdiction 以避免生成器陷入循环,例如,建议阅读文档。我目前也在研究字符概率,并提交了这个错误报告:github.com/huggingface/transformers/issues/16053。
    • @pete,你解决这个问题了吗?我需要同样的东西,从 generate() 获取每个标记的概率
    • 嗨@LearnToGrow 我刚刚发布了一个答案
    【解决方案2】:

    一个潜在的解决方法在线程https://github.com/huggingface/transformers/issues/10012 中。

    按照线程中的描述使用波束搜索,使用 n 个波束,其中 n 是您要显示的概率数,但只看未来的 1 个标记。然后,根据 mshuffett 的评论:

    我只是将此行移到 return_dict_in_generate 块下方。

    next_token_scores = next_token_scores + beam_scores[:, None].expand_as(next_token_scores)
    

    我试过了,效果很好。现在可以正确显示下一个标记的概率。

    或者,您可以尝试 https://github.com/huggingface/transformers/issues/16010 中描述的解决方案。我还没有解决它,因为它看起来比简单的解决方法稍微复杂一些。

    【讨论】:

    • 我不确定这段代码在做什么。我想要的是序列中与令牌对应的分数。这意味着通过对分数应用 softmax() 和 argmax(),我得到了 generate() 返回的相同序列索引。实际上,generate() 返回的是正确的分数。
    • 我不确定你的意思,我也不熟悉这段代码。我解决了原始问题中描述的问题:How to display the probabilities 1 token into the future。如果这不是您所期望的,那么您的问题可能有所不同。
    猜你喜欢
    • 2021-06-27
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2012-03-08
    • 1970-01-01
    • 2019-12-18
    • 2014-09-09
    • 2017-07-14
    相关资源
    最近更新 更多