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

使用NLTK的pad_both_ends生成n-grams时多余padding问题解决

问题:去除n-grams生成中多余的首尾padding

我正在编写函数为数据集中的每个短语生成n-grams,第一个短语示例如下。使用NLTK的pad_both_ends为短语添加<s>和</s> padding后,生成的bigrams首尾出现了额外两组padding,当前ans_n必须设为4,求去除方法。

原代码及问题输出

from nltk.lm.preprocessing import pad_both_ends
from nltk import bigrams

my_phrases[0] = ['It', 'is', 'my', 'favorite', 'place', 'ever', '.']

def my_ngrams(n, phrases):
    all_ngrams = []
    for i in range(len(phrases)):
        grams = list(bigrams(pad_both_ends(phrases[i], n)))
        all_ngrams.append(grams)
    return all_ngrams

ans_n = 4
ans_ngrams = my_ngrams(ans_n, my_phrases)
ans_ngrams

运行后输出(包含多余padding):

[('<s>', '<s>'),
 ('<s>', '<s>'),
 ('<s>', 'It'),
 ('It', 'is'),
 ('is', 'my'),
 ('my', 'favorite'),
 ('favorite', 'place'),
 ('place', 'ever'),
 ('ever', '.'),
 ('.', '</s>'),
 ('</s>', '</s>'),
 ('</s>', '</s>')]

期望输出

[('<s>', 'It'),
 ('It', 'is'),
 ('is', 'my'),
 ('my', 'favorite'),
 ('favorite', 'place'),
 ('place', 'ever'),
 ('ever', '.'),
 ('.', '</s>')]

问题原因

pad_both_ends的第二个参数是n-gram的阶数,当n=4时,它会在短语首尾各添加n-1=3个padding标记(即3个<s>前缀和3个</s>后缀)。直接对这个 padded 序列生成bigrams时,就会产生连续padding的无效bigrams(也就是首尾的两组重复padding)。

解决方法

方法1:过滤无效bigrams

直接过滤掉两个元素都是相同padding标记的bigrams,保留有效部分:

from nltk.lm.preprocessing import pad_both_ends
from nltk import bigrams

my_phrases = [['It', 'is', 'my', 'favorite', 'place', 'ever', '.']]

def my_ngrams(n, phrases):
    all_ngrams = []
    for phrase in phrases:
        padded_phrase = pad_both_ends(phrase, n)
        # 过滤掉两个元素都是padding的bigrams
        valid_grams = [gram for gram in bigrams(padded_phrase) if not (gram[0] == gram[1] and gram[0] in ('<s>', '</s>'))]
        all_ngrams.append(valid_grams)
    return all_ngrams

ans_n = 4
ans_ngrams = my_ngrams(ans_n, my_phrases)
print(ans_ngrams)

方法2:切片提取有效部分

因为当n=4时,首尾各有n-2=2个无效bigrams,直接通过切片去掉这些部分,效率更高:

from nltk.lm.preprocessing import pad_both_ends
from nltk import bigrams

my_phrases = [['It', 'is', 'my', 'favorite', 'place', 'ever', '.']]

def my_ngrams(n, phrases):
    all_ngrams = []
    for phrase in phrases:
        padded_phrase = pad_both_ends(phrase, n)
        # 切片去掉首尾各n-2个无效bigrams
        valid_grams = list(bigrams(padded_phrase))[n-2 : -(n-2)]
        all_ngrams.append(valid_grams)
    return all_ngrams

ans_n = 4
ans_ngrams = my_ngrams(ans_n, my_phrases)
print(ans_ngrams)

两种方法都能得到你想要的期望输出。

内容的提问来源于stack exchange,提问作者wellsinsideout

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 08:35:23