don't do td/bu prediction for epsilon categories if possible

This commit is contained in:
Krasimir Angelov
2026-09-04 16:25:11 +02:00
parent 51c376b9c2
commit a3b19e585d
3 changed files with 79 additions and 64 deletions
+17 -9
View File
@@ -1444,9 +1444,7 @@ void PgfParser::suspend(Cont *cont,Item *item,bool do_predict,ref<PgfSymbolCat>
auto &suspended = cont->suspended[value_i][lin_idx_i];
suspended.push_back(item);
size_t n_suspended = suspended.size();
if (cont->ccat == NULL) {
if (n_suspended == 1) {
if (suspended.size() == 1) {
std::function<void(ref<PgfCCat>)> f =
[this,item,cont](ref<PgfCCat> arg) {
@@ -1480,8 +1478,9 @@ void PgfParser::suspend(Cont *cont,Item *item,bool do_predict,ref<PgfSymbolCat>
cont->state->completed[cont][arg_ccat->value][arg_ccat->lin_idx] = arg_ccat;
};
epsilontable_iter(concr->epsilontable,cont->lincat,f);
}
if (cont->ccat == NULL) {
epsilontable_iter(concr->epsilontable,cont->lincat,0,f);
if (!cont->state->did_bu_predict) {
cont->state->did_bu_predict = true;
@@ -1492,20 +1491,28 @@ void PgfParser::suspend(Cont *cont,Item *item,bool do_predict,ref<PgfSymbolCat>
bu_predict(concr->phrasetable1, cont->state, viterbi_prob, 1, sentence->size);
}
}
} else {
if (do_predict && n_suspended == 1) {
if (cont->ccat->fid <= concr->last_fid) {
} else if (cont->ccat->fid <= concr->last_fid) {
epsilontable_iter(concr->epsilontable,cont->lincat,cont->ccat->fid,f);
if (do_predict) {
size_t n_items;
phrasetable_lookup(concr->phrasetable3,
cont->ccat->epsilon,
&n_items);
if (n_items == 0) {
for (size_t i = 0; i < cont->ccat->epsilon->n_items; i++) {
ref<PgfItem> pitem = cont->ccat->epsilon->items[i];
td_epsilon(cont->state,cont,pitem,item,symcat);
}
}
}
} else {
for (Production *prod : cont->ccat->prods) {
td_predict(cont->state,cont,prod,item,symcat);
}
}
}
}
State *state = cont->state;
while (state != NULL) {
@@ -1838,6 +1845,7 @@ void PgfParseTableMaker::final_item(State *state, CCat *ccat, Item *item, interv
epsilontable =
epsilontable_insert(epsilontable,
ccat->cont->lincat,
(ccat->cont->ccat != NULL) ? ccat->cont->ccat->fid : 0,
ccat->value, ccat->lin_idx,
ccat->fid, ccat->viterbi_prob,
pitem,
+13 -8
View File
@@ -576,7 +576,7 @@ vector<ref<PgfItem>> phrasetable_lookup(PgfPhrasetable<PgfSymbolBIND> phrasetabl
PGF_INTERNAL
PgfEpsilontable epsilontable_insert(PgfEpsilontable table,
ref<PgfConcrLincat> lincat,
ref<PgfConcrLincat> lincat, PgfMetaId prev_fid,
interval_t value, interval_t lin_idx,
PgfMetaId fid, prob_t viterbi_prob,
ref<PgfItem> item,
@@ -587,6 +587,7 @@ PgfEpsilontable epsilontable_insert(PgfEpsilontable table,
items[0] = item;
PgfEpsilontable new_table =
Node<PgfCCat>::new_node({.lincat=lincat,
.prev_fid=prev_fid,
.fid=fid,
.value=value,
.lin_idx=lin_idx,
@@ -604,12 +605,12 @@ PgfEpsilontable epsilontable_insert(PgfEpsilontable table,
if (cmp < 0) {
PgfEpsilontable left = epsilontable_insert(table->left,
lincat, value, lin_idx, fid, viterbi_prob, item, pepsilon);
lincat, prev_fid, value, lin_idx, fid, viterbi_prob, item, pepsilon);
table = Node<PgfCCat>::upd_node(table,left,table->right);
return Node<PgfCCat>::balanceL(table);
} else if (cmp > 0) {
PgfEpsilontable right = epsilontable_insert(table->right,
lincat, value, lin_idx, fid, viterbi_prob, item, pepsilon);
lincat, prev_fid, value, lin_idx, fid, viterbi_prob, item, pepsilon);
table = Node<PgfCCat>::upd_node(table, table->left, right);
return Node<PgfCCat>::balanceR(table);
} else {
@@ -665,20 +666,24 @@ ref<PgfCCat> epsilontable_get(PgfEpsilontable table,
}
PGF_INTERNAL
void epsilontable_iter(PgfEpsilontable table, ref<PgfConcrLincat> lincat, std::function<void(ref<PgfCCat> arg)> &f)
void epsilontable_iter(PgfEpsilontable table,
ref<PgfConcrLincat> lincat, PgfMetaId prev_fid,
std::function<void(ref<PgfCCat> arg)> &f)
{
if (table == 0)
return;
int cmp = textcmp(&lincat->name, &table->value.lincat->name);
if (cmp < 0)
epsilontable_iter(table->left, lincat, f);
epsilontable_iter(table->left, lincat, prev_fid, f);
else if (cmp > 0)
epsilontable_iter(table->right, lincat, f);
epsilontable_iter(table->right, lincat, prev_fid, f);
else {
epsilontable_iter(table->left, lincat, f);
epsilontable_iter(table->left, lincat, prev_fid, f);
if (table->value.prev_fid == prev_fid) {
f(ref<PgfCCat>::from_ptr(&table->value));
epsilontable_iter(table->right, lincat, f);
}
epsilontable_iter(table->right, lincat, prev_fid, f);
}
}
+5 -3
View File
@@ -50,7 +50,7 @@ struct PGF_INTERNAL_DECL PgfItem {
struct PGF_INTERNAL_DECL PgfCCat {
ref<PgfConcrLincat> lincat;
PgfMetaId fid;
PgfMetaId prev_fid, fid;
interval_t value, lin_idx;
prob_t viterbi_prob;
@@ -127,7 +127,7 @@ typedef ref<Node<PgfCCat>> PgfEpsilontable;
// The new category is mutable within the current transaction
PGF_INTERNAL_DECL
PgfEpsilontable epsilontable_insert(PgfEpsilontable table,
ref<PgfConcrLincat> lincat,
ref<PgfConcrLincat> lincat, PgfMetaId prev_fid,
interval_t value, interval_t lin_idx,
PgfMetaId fid, prob_t viterbi_prob,
ref<PgfItem> item,
@@ -143,7 +143,9 @@ ref<PgfCCat> epsilontable_get(PgfEpsilontable table,
PgfText *name, PgfMetaId fid);
PGF_INTERNAL
void epsilontable_iter(PgfEpsilontable table, ref<PgfConcrLincat> lincat, std::function<void(ref<PgfCCat> arg)> &f);
void epsilontable_iter(PgfEpsilontable table,
ref<PgfConcrLincat> lincat, PgfMetaId prev_fid,
std::function<void(ref<PgfCCat> arg)> &f);
PGF_INTERNAL_DECL
void epsilontable_release(PgfEpsilontable table);