Dev*_*per 7 python documentation generator huggingface-transformers
我正在使用 python Huggingfacetransformers库作为text-generation模型。我需要知道如何stopping_criteria在generator()我正在使用的函数中实现参数。
我stopping_criteria在本文档中找到了参数:
https://huggingface.co/transformers/main_classes/pipelines.html#transformers.TextGenerationPipeline
问题是,我只是不知道如何实施。
我的代码:
from transformers import pipeline
generator = pipeline('text-generation', model='EleutherAI/gpt-neo-125M')
stl = StoppingCriteria(['###'])
res = generator(prompt, do_sample=True,stopping_criteria = stl)
Run Code Online (Sandbox Code Playgroud)
小智 2
这两种方法对我有用。your_condition当你想停止时为 True。
class CustomStoppingCriteria(StoppingCriteria):
def __init__(self):
pass
def __call__(self, input_ids: torch.LongTensor, score: torch.FloatTensor, **kwargs) -> bool:
return your_condition
stopping_criteria = StoppingCriteriaList([CustomStoppingCriteria()])
Run Code Online (Sandbox Code Playgroud)
或者
def custom_stopping_criteria(input_ids: torch.LongTensor, score: torch.FloatTensor, **kwargs) -> bool:
return your_condition
stopping_criteria = StoppingCriteriaList([custom_stopping_criteria])
Run Code Online (Sandbox Code Playgroud)