Fairseq WMT19机器翻译模型.generate()函数返回值咨询
Fairseq
generate() 函数返回值说明 你通过torch.hub加载的transformer.wmt19.en-de模型,其generate()方法的返回值是嵌套列表结构:
- 外层列表对应输入的每个样本(比如输入N句话,外层列表就有N个元素)
- 内层列表包含该样本的所有生成候选(默认仅1个,若设置
beam参数会返回多个,按置信度排序)
每个候选是一个字典,核心字段如下:
'tokens': 生成的目标语言token张量(torch.Tensor类型),可通过模型的decode()方法转换为可读文本'score': 该生成序列的对数概率得分(float类型),数值越高代表模型对该序列的置信度越高'attention': 可选字段,仅当生成时指定attention=True时返回,包含编码器-解码器的注意力权重'alignment': 可选字段,仅当生成时指定print_alignment=True时返回,包含源语言与目标语言token的对齐信息
示例代码
# 单句输入示例 output = en2de.generate("Hello, how are you?") # 提取第一个样本的第一个候选文本 translated_text = en2de.decode(output[0][0]['tokens']) print(translated_text)
如果设置beam=3开启beam搜索,每个输入样本的内层列表会包含3个候选字典,按score从高到低排列。
内容的提问来源于stack exchange,提问作者Deergh Singh
相关产品推荐
相关产品推荐

