【问题标题】:Summarization with Huggingface: How to generate one word at a time?Huggingface 总结:如何一次生成一个单词?
【发布时间】:2022-06-30 19:32:08
【问题描述】:

我正在使用 DistilBART 进行抽象摘要。 generate() 方法使用起来非常简单。但是,它会返回完整的、完成的摘要。 我想要的是,在每个步骤中,访问 logits,然后获取下一个单词候选列表并根据我自己的标准进行选择。一旦选择,继续下一个单词,依此类推,直到生成 EOS 代币。

我知道我可以通过 model(**input).logits[:, -1, :] 访问 logits,但这里的输入将是整个(编码的)文本,那么这些 logits 到底对应的是什么?第一个生成的令牌?最后一个?

感谢您的回答!

【问题讨论】:

标签: huggingface-transformers summarization huggingface


【解决方案1】:

为了将来参考,这是如何完成的注意:这是特定于编码器-解码器模型,如 BART):

1.初始化

import torch
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM

# Load model
tokenizer = AutoTokenizer.from_pretrained("sshleifer/distilbart-xsum-1-1")
model = AutoModelForSeq2SeqLM.from_pretrained("sshleifer/distilbart-xsum-1-1")

text = "..."

# Tokenize text
batch = tokenizer(text, return_tensors="pt")

2.选项 1:使用贪婪解码(无缓存)生成摘要

generated_sequence = torch.tensor([[tokenizer.sep_token_id]])  # initial token

# Generation loop
while True:
    with torch.no_grad():
        output = model(input_ids=batch["input_ids"], decoder_input_ids=generated_sequence)
    next_token_logits = output.logits[:, -1, :]
    next_token_scores = next_token_logits.softmax(dim=-1)

    # Take token with highest probability
    next_token = next_token_scores.argmax().unsqueeze(0).unsqueeze(0)

    # Append token to generated sequence
    generated_sequence = torch.cat((generated_sequence, next_token), dim=1)
    # Stop if EOS token generated
    if (generated_sequence.squeeze()[-1] == tokenizer.eos_token_id):
        break

summary = tokenizer.batch_decode(generated_sequence, skip_special_tokens=True)

3.选项 2:使用 top-k、top-p 采样和温度生成摘要(无缓存)

from transformers.generation_utils import top_k_top_p_filtering

generated_sequence = torch.tensor([[tokenizer.sep_token_id]])  # initial token

# Generation loop
while True:
    with torch.no_grad():
        output = model(input_ids=batch["input_ids"], decoder_input_ids=generated_sequence)
    logits = output.logits[:, -1, :] / temperature  # apply temperature
    filtered_logits = top_k_top_p_filtering(logits=logits, top_k=4, top_p=0.7)
    probabilities = filtered_logits.softmax(dim=-1)

    # Sample next token
    next_token = torch.multinomial(probabilities, 1)

    # Append token to generated sequence
    generated_sequence = torch.cat((generated_sequence, next_token), dim=1)
    # Stop if EOS token generated
    if (generated_sequence.squeeze()[-1] == tokenizer.eos_token_id):
        break

summary = tokenizer.batch_decode(generated_sequence, skip_special_tokens=True)

(其他 generating strategies 类似)。


4.使用缓存

由于编码器的输入(即要总结的文本)总是相同的,我们可以将其缓存以大大加快生成速度。

generated_sequence = torch.tensor([[tokenizer.sep_token_id]])  # initial token
input_ids = batch["input_ids"]
past_key_values = None

with torch.no_grad():
    output = model(
        input_ids=input_ids,
        decoder_input_ids=generated_sequence,
        past_key_values=past_key_values
    )
    
encoder_outputs=output.encoder_last_hidden_state

# Generation loop
while True:
    # From here on, use cached attention
    past_key_values = output.past_key_values
    next_token_logits = output.logits[:, -1, :]
    next_token_scores = next_token_logits.softmax(dim=-1)
    next_token = next_token_scores.argmax().unsqueeze(0).unsqueeze(0)  # greedy decoding
    generated_sequence = torch.cat((generated_sequence, next_token), dim=1)
    # Stop if EOS token generated
    if (generated_sequence.squeeze()[-1] == tokenizer.eos_token_id):
        break
    with torch.no_grad():
        output = model(
            decoder_input_ids=torch.tensor([[generated_sequence.squeeze()[-1]]]),
            past_key_values=past_key_values,
            encoder_outputs=encoder_outputs
        )

summary = tokenizer.batch_decode(generated_sequence, skip_special_tokens=True)

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 2014-06-03
    • 1970-01-01
    • 2022-01-14
    • 1970-01-01
    • 2014-02-03
    • 1970-01-01
    相关资源
    最近更新 更多