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

如何修改Python的init_wfst与complete_wfst函数实现WFST单元格存储非终结符集

WFST单元格支持存储非终结符集合的修改方案

核心调整逻辑如下:

  • 将原WFST单元格存储单个非终结符/None的逻辑,改为存储非终结符集合,空集合代表当前跨度无匹配的非终结符
  • 语法规则索引从「右部对应单个左部」改为「右部对应左部集合」,支持同一种右部匹配多条产生式
  • 规则匹配阶段遍历两个中间单元格所有非终结符的组合,所有符合规则的结果都存入目标单元格

1 修改后的init_wfst函数

初始化时每个单元格默认是空集合,将所有匹配当前词的产生式左部全部加入对应单元格:

def init_wfst(tokens, grammar):
    numtokens = len(tokens)
    # 初始化所有单元格为空白非终结符集合
    wfst = [[set() for i in range(numtokens + 1)] for j in range(numtokens+1)]
    for i in range(numtokens):
        productions = grammar.productions(rhs=tokens[i])
        # 将所有匹配的非终结符全部加入集合,不再只取第一个
        for p in productions:
            wfst[i][i+1].add(p.lhs())
    return wfst

2 修改后的complete_wfst函数

调整语法索引结构,遍历所有非终结符组合匹配规则:

def complete_wfst(wfst, tokens, grammar, trace=True):
    # 构建索引:同一个右部可能对应多个左部,所以值用集合存储
    index = {}
    for p in grammar.productions():
        rhs = p.rhs()
        if rhs not in index:
            index[rhs] = set()
        index[rhs].add(p.lhs())
    numtokens = len(tokens)
    for span in range(2, numtokens + 1):
        for start in range(numtokens + 1 - span):
            end = start + span
            for mid in range(start + 1, end):
                nt1_set, nt2_set = wfst[start][mid], wfst[mid][end]
                # 遍历两个集合的所有非终结符组合
                for nt1 in nt1_set:
                    for nt2 in nt2_set:
                        key = (nt1, nt2)
                        if key in index:
                            # 将所有匹配的左部加入目标单元格
                            for lhs in index[key]:
                                wfst[start][end].add(lhs)
                                if trace:
                                    print("[%s] %3s [%s] %3s [%s] ==> [%s] %3s [%s]" %
                                          (start, nt1, mid, nt2, end, start, lhs, end))
    return wfst

3 修改后的display函数

适配集合类型的展示逻辑,多个非终结符用逗号分隔:

def display(wfst, tokens):
    print('\nWFST ' + ' '.join(("%-4d" % i) for i in range(1, len(wfst))))
    for i in range(len(wfst) - 1):
        print("%d   " % i, end=" ")
        for j in range(1, len(wfst)):
            cell_val = wfst[i][j]
            # 空集合显示. 非空则拼接所有非终结符
            print_str = ','.join(str(nt) for nt in cell_val) if cell_val else '.'
            print("%-4s" % print_str, end=" ")
        print()

修改后的代码调用逻辑和原有逻辑完全一致,不需要调整:

tokens = "I shot an elephant in my pajamas".split()
wfst0 = init_wfst(tokens, groucho_grammar)
display(wfst0, tokens)

wfst1 = complete_wfst(wfst0, tokens, groucho_grammar)
display(wfst1, tokens)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.24 12:45:03