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]; auto &suspended = cont->suspended[value_i][lin_idx_i];
suspended.push_back(item); suspended.push_back(item);
size_t n_suspended = suspended.size(); if (suspended.size() == 1) {
if (cont->ccat == NULL) {
if (n_suspended == 1) {
std::function<void(ref<PgfCCat>)> f = std::function<void(ref<PgfCCat>)> f =
[this,item,cont](ref<PgfCCat> arg) { [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; 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) { if (!cont->state->did_bu_predict) {
cont->state->did_bu_predict = true; 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); bu_predict(concr->phrasetable1, cont->state, viterbi_prob, 1, sentence->size);
} }
} }
} else { } else if (cont->ccat->fid <= concr->last_fid) {
if (do_predict && n_suspended == 1) { epsilontable_iter(concr->epsilontable,cont->lincat,cont->ccat->fid,f);
if (cont->ccat->fid <= concr->last_fid) {
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++) { for (size_t i = 0; i < cont->ccat->epsilon->n_items; i++) {
ref<PgfItem> pitem = cont->ccat->epsilon->items[i]; ref<PgfItem> pitem = cont->ccat->epsilon->items[i];
td_epsilon(cont->state,cont,pitem,item,symcat); td_epsilon(cont->state,cont,pitem,item,symcat);
} }
}
}
} else { } else {
for (Production *prod : cont->ccat->prods) { for (Production *prod : cont->ccat->prods) {
td_predict(cont->state,cont,prod,item,symcat); td_predict(cont->state,cont,prod,item,symcat);
} }
} }
} }
}
State *state = cont->state; State *state = cont->state;
while (state != NULL) { while (state != NULL) {
@@ -1838,6 +1845,7 @@ void PgfParseTableMaker::final_item(State *state, CCat *ccat, Item *item, interv
epsilontable = epsilontable =
epsilontable_insert(epsilontable, epsilontable_insert(epsilontable,
ccat->cont->lincat, ccat->cont->lincat,
(ccat->cont->ccat != NULL) ? ccat->cont->ccat->fid : 0,
ccat->value, ccat->lin_idx, ccat->value, ccat->lin_idx,
ccat->fid, ccat->viterbi_prob, ccat->fid, ccat->viterbi_prob,
pitem, pitem,
+13 -8
View File
@@ -576,7 +576,7 @@ vector<ref<PgfItem>> phrasetable_lookup(PgfPhrasetable<PgfSymbolBIND> phrasetabl
PGF_INTERNAL PGF_INTERNAL
PgfEpsilontable epsilontable_insert(PgfEpsilontable table, PgfEpsilontable epsilontable_insert(PgfEpsilontable table,
ref<PgfConcrLincat> lincat, ref<PgfConcrLincat> lincat, PgfMetaId prev_fid,
interval_t value, interval_t lin_idx, interval_t value, interval_t lin_idx,
PgfMetaId fid, prob_t viterbi_prob, PgfMetaId fid, prob_t viterbi_prob,
ref<PgfItem> item, ref<PgfItem> item,
@@ -587,6 +587,7 @@ PgfEpsilontable epsilontable_insert(PgfEpsilontable table,
items[0] = item; items[0] = item;
PgfEpsilontable new_table = PgfEpsilontable new_table =
Node<PgfCCat>::new_node({.lincat=lincat, Node<PgfCCat>::new_node({.lincat=lincat,
.prev_fid=prev_fid,
.fid=fid, .fid=fid,
.value=value, .value=value,
.lin_idx=lin_idx, .lin_idx=lin_idx,
@@ -604,12 +605,12 @@ PgfEpsilontable epsilontable_insert(PgfEpsilontable table,
if (cmp < 0) { if (cmp < 0) {
PgfEpsilontable left = epsilontable_insert(table->left, 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); table = Node<PgfCCat>::upd_node(table,left,table->right);
return Node<PgfCCat>::balanceL(table); return Node<PgfCCat>::balanceL(table);
} else if (cmp > 0) { } else if (cmp > 0) {
PgfEpsilontable right = epsilontable_insert(table->right, 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); table = Node<PgfCCat>::upd_node(table, table->left, right);
return Node<PgfCCat>::balanceR(table); return Node<PgfCCat>::balanceR(table);
} else { } else {
@@ -665,20 +666,24 @@ ref<PgfCCat> epsilontable_get(PgfEpsilontable table,
} }
PGF_INTERNAL 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) if (table == 0)
return; return;
int cmp = textcmp(&lincat->name, &table->value.lincat->name); int cmp = textcmp(&lincat->name, &table->value.lincat->name);
if (cmp < 0) if (cmp < 0)
epsilontable_iter(table->left, lincat, f); epsilontable_iter(table->left, lincat, prev_fid, f);
else if (cmp > 0) else if (cmp > 0)
epsilontable_iter(table->right, lincat, f); epsilontable_iter(table->right, lincat, prev_fid, f);
else { 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)); 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 { struct PGF_INTERNAL_DECL PgfCCat {
ref<PgfConcrLincat> lincat; ref<PgfConcrLincat> lincat;
PgfMetaId fid; PgfMetaId prev_fid, fid;
interval_t value, lin_idx; interval_t value, lin_idx;
prob_t viterbi_prob; prob_t viterbi_prob;
@@ -127,7 +127,7 @@ typedef ref<Node<PgfCCat>> PgfEpsilontable;
// The new category is mutable within the current transaction // The new category is mutable within the current transaction
PGF_INTERNAL_DECL PGF_INTERNAL_DECL
PgfEpsilontable epsilontable_insert(PgfEpsilontable table, PgfEpsilontable epsilontable_insert(PgfEpsilontable table,
ref<PgfConcrLincat> lincat, ref<PgfConcrLincat> lincat, PgfMetaId prev_fid,
interval_t value, interval_t lin_idx, interval_t value, interval_t lin_idx,
PgfMetaId fid, prob_t viterbi_prob, PgfMetaId fid, prob_t viterbi_prob,
ref<PgfItem> item, ref<PgfItem> item,
@@ -143,7 +143,9 @@ ref<PgfCCat> epsilontable_get(PgfEpsilontable table,
PgfText *name, PgfMetaId fid); PgfText *name, PgfMetaId fid);
PGF_INTERNAL 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 PGF_INTERNAL_DECL
void epsilontable_release(PgfEpsilontable table); void epsilontable_release(PgfEpsilontable table);