minimize top-down predictions

This commit is contained in:
Krasimir Angelov
2026-08-17 14:33:14 +02:00
parent 488b424626
commit b7a911faf8
2 changed files with 186 additions and 177 deletions
+175 -170
View File
@@ -150,16 +150,14 @@ void PgfAbstractParser::symbol(Item *item, State *state, PgfSymbol sym)
cont->state = state;
}
interval_t value_i = item->interval(item->rule->args[symcat->d]);
interval_t lin_idx_i = item->interval(ref<PgfLParam>::from_ptr(&symcat->r));
auto &suspended = cont->suspended[value_i][lin_idx_i];
suspended.push_back(item);
interval_t value_i = interval(item->rule, &item->vars[0], item->rule->args[symcat->d]);
interval_t lin_idx_i = interval(item->rule, &item->vars[0], ref<PgfLParam>::from_ptr(&symcat->r));
suspend(cont,item,n_suspended1 == 0,suspended.size(),symcat);
suspend(cont,item,n_suspended1 == 0,symcat,value_i,lin_idx_i);
}
} else {
interval_t value_i = item->interval(item->rule->args[symcat->d]);
interval_t lin_idx_i = item->interval(ref<PgfLParam>::from_ptr(&symcat->r));
interval_t value_i = interval(item->rule, &item->vars[0], item->rule->args[symcat->d]);
interval_t lin_idx_i = interval(item->rule, &item->vars[0], ref<PgfLParam>::from_ptr(&symcat->r));
// the following prevents infinite loops with epsilons
bool found = false;
@@ -199,12 +197,8 @@ void PgfAbstractParser::symbol(Item *item, State *state, PgfSymbol sym)
}
}
}
found:;
auto &suspended = cont->suspended[value_i][lin_idx_i];
suspended.push_back(item);
suspend(cont,item,!subsumed,suspended.size(),symcat);
found:
suspend(cont,item,!subsumed,symcat,value_i,lin_idx_i);
}
break;
}
@@ -264,8 +258,8 @@ void PgfAbstractParser::complete(Item *item, State *state)
case PgfConcrLin::tag: {
auto lin = ref<PgfConcrLin>::untagged(item->rule->container);
interval_t res = item->interval(item->rule->res);
interval_t lin_idx = item->interval(item->rule->lin_idx);
interval_t res = interval(item->rule, &item->vars[0], item->rule->res);
interval_t lin_idx = interval(item->rule, &item->vars[0], item->rule->lin_idx);
CCat *&ccat = state->completed[item->cont][res][lin_idx];
if (ccat == NULL) {
ccat = new CCat;
@@ -360,38 +354,25 @@ void PgfAbstractParser::complete(Item *item, State *state)
delete item;
}
interval_t PgfAbstractParser::Item::interval(ref<PgfLParam> lparam) const
{
interval_t interval;
interval.first = lparam->i0;
interval.second = interval.first;
for (size_t i = 0; i < lparam->n_terms; i++) {
size_t var = lparam->terms[i].var;
if (vars[var] == 0) {
interval.second += lparam->terms[i].factor * (rule->ranges[var]-1);
} else {
size_t value = lparam->terms[i].factor * (vars[var]-1);
interval.first += value;
interval.second += value;
}
}
return interval;
}
#define ZERO_VALUES(rule) \
((size_t*) memset(alloca(rule->ranges.size()*sizeof(size_t)), 0, rule->ranges.size()*sizeof(size_t)))
#define CLONE_VALUES(rule,values) \
((size_t*) memcpy(alloca(rule->ranges.size()*sizeof(size_t)), values, rule->ranges.size()*sizeof(size_t)))
bool PgfAbstractParser::Item::instantiate(ref<PgfLParam> lparam1,
PgfConcrRule *rule, size_t *values, ref<PgfLParam> lparam2)
bool PgfAbstractParser::instantiate(ref<PgfConcrRule> rule1, size_t *values1, ref<PgfLParam> lparam1,
ref<PgfConcrRule> rule2, size_t *values2, ref<PgfLParam> lparam2)
{
size_t i01 = lparam1->i0;
for (size_t i = 0; i < lparam1->n_terms; i++) {
if (this->vars[lparam1->terms[i].var] > 0) {
i01 += lparam1->terms[i].factor * (this->vars[lparam1->terms[i].var]-1);
if (values1[lparam1->terms[i].var] > 0) {
i01 += lparam1->terms[i].factor * (values1[lparam1->terms[i].var]-1);
}
}
size_t i02 = lparam2->i0;
for (size_t i = 0; i < lparam2->n_terms; i++) {
if (values[lparam2->terms[i].var] > 0) {
i02 += lparam2->terms[i].factor * (values[lparam2->terms[i].var]-1);
if (values2[lparam2->terms[i].var] > 0) {
i02 += lparam2->terms[i].factor * (values2[lparam2->terms[i].var]-1);
}
}
@@ -409,22 +390,22 @@ bool PgfAbstractParser::Item::instantiate(ref<PgfLParam> lparam1,
term t1 = {0,0};
if (i1 < lparam1->n_terms) {
t1 = lparam1->terms[i1];
if (this->vars[t1.var] > 0) {
if (values1[t1.var] > 0) {
i1++;
continue;
}
scale1 = t1.factor * this->rule->ranges[t1.var];
scale1 = t1.factor * rule1->ranges[t1.var];
}
size_t scale2 = 0;
term t2 = {0,0};
if (i2 < lparam2->n_terms) {
t2 = lparam2->terms[i2];
if (values[t2.var] > 0) {
if (values2[t2.var] > 0) {
i2++;
continue;
}
scale2 = t2.factor * rule->ranges[t2.var];
scale2 = t2.factor * rule2->ranges[t2.var];
}
if (scale1 > scale2) {
@@ -436,20 +417,20 @@ bool PgfAbstractParser::Item::instantiate(ref<PgfLParam> lparam1,
if (f == 0)
break;
if (values[t2.var] == 0) {
max += f * (rule->ranges[t2.var]-1);
if (values2[t2.var] == 0) {
max += f * (rule2->ranges[t2.var]-1);
}
i2++;
}
i02 %= t1.factor;
if (min >= this->rule->ranges[t1.var])
if (min >= rule1->ranges[t1.var])
return false;
if (min == max) {
if (this->vars[t1.var] == 0)
this->vars[t1.var] = min+1;
else if (this->vars[t1.var] != min+1)
if (values1[t1.var] == 0)
values1[t1.var] = min+1;
else if (values1[t1.var] != min+1)
return false;
}
@@ -463,30 +444,48 @@ bool PgfAbstractParser::Item::instantiate(ref<PgfLParam> lparam1,
if (f == 0)
break;
if (values[t1.var] == 0) {
max += f * (rule->ranges[t1.var]-1);
if (values1[t1.var] == 0) {
max += f * (rule1->ranges[t1.var]-1);
}
i1++;
}
i01 %= t2.factor;
if (min >= rule->ranges[t2.var])
if (min >= rule2->ranges[t2.var])
return false;
if (min == max) {
if (values[t2.var] == 0) {
// we don't update the production;
} else if (values[t2.var] != min+1)
if (values2[t2.var] == 0) {
values2[t2.var] = min+1;
} else if (values2[t2.var] != min+1)
return false;
}
i2++;
}
}
return (i01 == i02);
}
interval_t PgfAbstractParser::interval(ref<PgfConcrRule> rule, size_t *values, ref<PgfLParam> lparam)
{
interval_t interval;
interval.first = lparam->i0;
interval.second = interval.first;
for (size_t i = 0; i < lparam->n_terms; i++) {
size_t var = lparam->terms[i].var;
if (values[var] == 0) {
interval.second += lparam->terms[i].factor * (rule->ranges[var]-1);
} else {
size_t value = lparam->terms[i].factor * (values[var]-1);
interval.first += value;
interval.second += value;
}
}
return interval;
}
void PgfAbstractParser::combine(State *state, Item *item, CCat *ccat)
{
PgfSymbol sym = item->rule->syms[item->dot];
@@ -495,11 +494,15 @@ void PgfAbstractParser::combine(State *state, Item *item, CCat *ccat)
ref<PgfConcrRule> rule;
size_t *values;
get_info(ccat, &rule,&values);
if (!item->instantiate(item->rule->args[sym_cat->d], rule, values, rule->res)) {
values = CLONE_VALUES(rule, values);
if (!instantiate(item->rule, &item->vars[0], item->rule->args[sym_cat->d],
rule, values, rule->res)) {
delete item;
return;
}
if (!item->instantiate(ref<PgfLParam>::from_ptr(&sym_cat->r), rule, values, rule->lin_idx)) {
if (!instantiate(item->rule, &item->vars[0], ref<PgfLParam>::from_ptr(&sym_cat->r),
rule, values, rule->lin_idx)) {
delete item;
return;
}
@@ -520,38 +523,52 @@ void PgfAbstractParser::td_epsilon(State *state, Cont *cont, ref<PgfItem> pitem,
auto lin = ref<PgfConcrLin>::untagged(pitem->rule->container);
for (ref<PgfConcrRule> rule : lin->rules) {
Item *item = new (rule) Item;
item->cont = cont;
item->dot = 0;
item->pre_alt = 0;
item->pre_dot = 0;
item->syms = rule->syms.as_vector();
item->rule = rule;
item->inside_prob = lin->absfun->prob;
item->outside_prob = xitem->outside_prob+xitem->inside_prob-xitem->args[symcat->d]->viterbi_prob;
if (!item->instantiate(item->rule->res, xitem->rule, &xitem->vars[0], xitem->rule->args[symcat->d])) {
delete item;
size_t *values1 = ZERO_VALUES(rule);
size_t *values2 = CLONE_VALUES(xitem->rule, &xitem->vars[0]);
if (!instantiate(rule, values1, rule->res,
xitem->rule, values2, xitem->rule->args[symcat->d])) {
continue;
}
if (!item->instantiate(item->rule->lin_idx, xitem->rule, &xitem->vars[0], ref<PgfLParam>::from_ptr(&symcat->r))) {
delete item;
if (!instantiate(rule, values1, rule->lin_idx,
xitem->rule, values2, ref<PgfLParam>::from_ptr(&symcat->r))) {
continue;
}
size_t *values3 = CLONE_VALUES(pitem->rule, &pitem->vars[0]);
for (size_t i = 0; i < pitem->args.size(); i++) {
if (pitem->args[i] != 0) {
item->args[i] = get_epsilon_ccat(&lin->absfun->type->hypos[i].type->name,pitem->args[i]);
item->inside_prob += item->args[i]->viterbi_prob;
}
if (!item->instantiate(item->rule->args[i], pitem->rule, &pitem->vars[0], pitem->rule->args[i])) {
delete item;
if (!instantiate(rule, values1, rule->args[i],
pitem->rule, values3, pitem->rule->args[i])) {
goto next;
}
}
state->push_item(item);
{
interval_t value_i = interval(rule, values1, rule->res);
interval_t lin_idx_i = interval(rule, values1, rule->lin_idx);
Item *&pred = cont->predicted[value_i][lin_idx_i];
if (pred != NULL && pred != xitem)
return;
pred = xitem;
Item *item = new (rule) Item;
item->cont = cont;
item->dot = 0;
item->pre_alt = 0;
item->pre_dot = 0;
item->syms = rule->syms.as_vector();
item->rule = rule;
item->inside_prob = lin->absfun->prob;
item->outside_prob = xitem->outside_prob+xitem->inside_prob-xitem->args[symcat->d]->viterbi_prob;
for (size_t i = 0; i < pitem->args.size(); i++) {
if (pitem->args[i] != 0) {
item->args[i] = get_epsilon_ccat(&lin->absfun->type->hypos[i].type->name,pitem->args[i]);
item->inside_prob += item->args[i]->viterbi_prob;
}
}
state->push_item(item);
}
next:;
}
}
@@ -567,38 +584,54 @@ void PgfAbstractParser::td_predict(State *state, Cont *cont, Production *prod, I
auto lin = ref<PgfConcrLin>::untagged(prod->rule->container);
for (ref<PgfConcrRule> rule : lin->rules) {
Item *item = new (rule) Item;
item->cont = cont;
item->dot = 0;
item->pre_alt = 0;
item->pre_dot = 0;
item->syms = rule->syms.as_vector();
item->rule = rule;
item->inside_prob = lin->absfun->prob;
item->outside_prob = xitem->outside_prob+xitem->inside_prob-xitem->args[symcat->d]->viterbi_prob;
if (!item->instantiate(item->rule->res, xitem->rule, &xitem->vars[0], xitem->rule->args[symcat->d])) {
delete item;
size_t *values1 = ZERO_VALUES(rule);
size_t *values2 = CLONE_VALUES(xitem->rule, &xitem->vars[0]);
if (!instantiate(rule, values1, rule->res,
xitem->rule, values2, xitem->rule->args[symcat->d])) {
continue;
}
if (!item->instantiate(item->rule->lin_idx, xitem->rule, &xitem->vars[0], ref<PgfLParam>::from_ptr(&symcat->r))) {
delete item;
if (!instantiate(rule, values1, rule->lin_idx,
xitem->rule, values2, ref<PgfLParam>::from_ptr(&symcat->r))) {
continue;
}
for (size_t i = 0; i < item->args.size(); i++) {
if (!item->instantiate(item->rule->args[i], prod->rule, &prod->vars[0], prod->rule->args[i])) {
delete item;
size_t *values3 = CLONE_VALUES(prod->rule, &prod->vars[0]);
for (size_t i = 0; i < rule->args.size(); i++) {
if (!instantiate(rule, values1, rule->args[i],
prod->rule, values3, prod->rule->args[i])) {
goto next;
}
item->args[i] = prod->args[i];
if (item->args[i] != NULL) {
item->inside_prob += item->args[i]->viterbi_prob;
}
}
state->push_item(item);
{
interval_t value_i = interval(rule, values1, rule->res);
interval_t lin_idx_i = interval(rule, values1, rule->lin_idx);
Item *&pred = cont->predicted[value_i][lin_idx_i];
if (pred != NULL && pred != xitem)
return;
pred = xitem;
Item *item = new (rule) Item;
item->cont = cont;
item->dot = 0;
item->pre_alt = 0;
item->pre_dot = 0;
item->syms = rule->syms.as_vector();
item->rule = rule;
item->inside_prob = lin->absfun->prob;
item->outside_prob = xitem->outside_prob+xitem->inside_prob-xitem->args[symcat->d]->viterbi_prob;
for (size_t i = 0; i < rule->args.size(); i++) {
item->args[i] = prod->args[i];
if (item->args[i] != NULL) {
item->inside_prob += item->args[i]->viterbi_prob;
}
}
state->push_item(item);
}
next:;
}
}
@@ -1364,24 +1397,29 @@ void PgfParser::symbol_bind(Item *item, State *state, PgfSymbol sym)
}
}
void PgfParser::suspend(Cont *cont,Item *item,bool do_predict,size_t n_suspended,ref<PgfSymbolCat> symcat)
void PgfParser::suspend(Cont *cont,Item *item,bool do_predict,ref<PgfSymbolCat> symcat,interval_t value_i,interval_t lin_idx_i)
{
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) {
std::function<void(ref<PgfCCat>)> f =
[this,item,cont](ref<PgfCCat> arg) {
ref<PgfItem> xitem = arg->items[0];
ref<PgfItem> pitem = arg->items[0];
Item *new_item = new (item) Item;
PgfSymbol sym = new_item->rule->syms[new_item->dot];
PgfSymbol sym = item->rule->syms[item->dot];
auto sym_cat = ref<PgfSymbolCat>::untagged(sym);
if (!new_item->instantiate(new_item->rule->args[sym_cat->d],xitem->rule,&xitem->vars[0],xitem->rule->res)) {
delete new_item;
size_t *values1 = CLONE_VALUES(item->rule, &item->vars[0]);
size_t *values2 = CLONE_VALUES(pitem->rule, &pitem->vars[0]);
if (!instantiate(item->rule, values1, item->rule->args[sym_cat->d],
pitem->rule, values2, pitem->rule->res)) {
return;
}
if (!new_item->instantiate(ref<PgfLParam>::from_ptr(&sym_cat->r),xitem->rule,&xitem->vars[0],xitem->rule->lin_idx)) {
delete new_item;
if (!instantiate(item->rule, values1, ref<PgfLParam>::from_ptr(&sym_cat->r),
pitem->rule, values2, pitem->rule->lin_idx)) {
return;
}
@@ -1399,30 +1437,10 @@ void PgfParser::suspend(Cont *cont,Item *item,bool do_predict,size_t n_suspended
}
cont->state->completed[cont][arg_ccat->value][arg_ccat->lin_idx] = arg_ccat;
new_item->dot++;
new_item->args[sym_cat->d] = arg_ccat;
new_item->inside_prob += arg_ccat->viterbi_prob;
cont->state->push_item(new_item);
};
epsilontable_iter(concr->epsilontable,cont->lincat,f);
}
State *state = cont->state;
while (state != NULL) {
auto it1 = state->completed.find(cont);
if (it1 != state->completed.end()) {
for (auto it2 : it1->second) {
for (auto it3 : it2.second) {
Item *new_item = new (item) Item;
combine(state, new_item, it3.second);
}
}
}
state = state->next;
}
if (do_predict) {
if (cont->state->needs_bind) {
bu_predict(concr->phrasetable4, cont->state);
@@ -1442,23 +1460,21 @@ void PgfParser::suspend(Cont *cont,Item *item,bool do_predict,size_t n_suspended
td_predict(cont->state,cont,prod,item,symcat);
}
}
} else {
interval_t lin_idx_i = item->interval(ref<PgfLParam>::from_ptr(&symcat->r));
State *next = cont->state;
while (next != NULL) {
auto it1 = next->completed.find(cont);
if (it1 != next->completed.end()) {
for (auto it2 : it1->second.overlaps(cont->ccat->value)) {
for (auto it3 : it2.second.overlaps(lin_idx_i)) {
CCat *arg = it3.second;
Item *new_item = new (item) Item;
combine(next, new_item, arg);
}
}
}
}
State *state = cont->state;
while (state != NULL) {
auto it1 = state->completed.find(cont);
if (it1 != state->completed.end()) {
for (auto it2 : it1->second.overlaps(value_i)) {
for (auto it3 : it2.second.overlaps(lin_idx_i)) {
Item *new_item = new (item) Item;
combine(state, new_item, it3.second);
}
next = next->next;
}
}
state = state->next;
}
}
@@ -1610,19 +1626,13 @@ void PgfParseTableMaker::symbol_bind(Item *item, State *state, PgfSymbol sym)
}
}
void PgfParseTableMaker::suspend(Cont *cont,Item *item,bool do_predict,size_t n_suspended,ref<PgfSymbolCat> symcat)
void PgfParseTableMaker::suspend(Cont *cont,Item *item,bool do_predict,ref<PgfSymbolCat> symcat,interval_t value_i,interval_t lin_idx_i)
{
if (cont->ccat == NULL) {
for (auto it1 : cont->state->completed[cont]) {
for (auto it2 : it1.second) {
CCat *ccat = it2.second;
if (ccat != NULL) {
Item *new_item = new (item) Item;
combine(cont->state,new_item,ccat);
}
}
}
auto &suspended = cont->suspended[value_i][lin_idx_i];
suspended.push_back(item);
size_t n_suspended = suspended.size();
if (cont->ccat == NULL) {
auto pitem = clone_item(item);
auto phrasetable2 = phrasetable_insert(concr->phrasetable2,cont->lincat,pitem);
concr->phrasetable2 = phrasetable2;
@@ -1638,28 +1648,23 @@ void PgfParseTableMaker::suspend(Cont *cont,Item *item,bool do_predict,size_t n_
td_predict(cont->state,cont,prod,item,symcat);
}
}
} else {
interval_t lin_idx_i = item->interval(ref<PgfLParam>::from_ptr(&symcat->r));
State *next = cont->state;
while (next != NULL) {
auto it1 = next->completed.find(cont);
if (it1 != next->completed.end()) {
for (auto it2 : it1->second.overlaps(cont->ccat->value)) {
for (auto it3 : it2.second.overlaps(lin_idx_i)) {
CCat *arg = it3.second;
Item *new_item = new (item) Item;
combine(next, new_item, arg);
}
}
}
next = next->next;
}
}
auto pitem = clone_item(item);
auto phrasetable3 = phrasetable_insert(concr->phrasetable3,cont->ccat->epsilon,pitem);
concr->phrasetable3 = phrasetable3;
}
auto it1 = cont->state->completed.find(cont);
if (it1 != cont->state->completed.end()) {
for (auto it2 : it1->second.overlaps(value_i)) {
for (auto it3 : it2.second.overlaps(lin_idx_i)) {
CCat *arg = it3.second;
Item *new_item = new (item) Item;
combine(cont->state, new_item, arg);
}
}
}
}
void PgfParseTableMaker::final_item(State *state, CCat *ccat, Item *item, interval_t value, interval_t lin_idx)
+11 -7
View File
@@ -128,6 +128,7 @@ protected:
ref<PgfConcrLincat> lincat;
State *state;
interval_map<interval_map<std::vector<Item*>>> suspended;
interval_map<interval_map<Item*>> predicted;
~Cont();
};
@@ -189,10 +190,6 @@ protected:
Item() {
}
interval_t interval(ref<PgfLParam> lparam) const;
bool instantiate(ref<PgfLParam> lparam1,
PgfConcrRule *rule, size_t *values, ref<PgfLParam> lparam2);
};
static struct ItemComparator : std::less<Item*> {
@@ -239,7 +236,7 @@ protected:
virtual State *new_state(const PgfTextSpot &start)=0;
virtual void symbol_token(Item *item, State *state, ref<PgfSymbolKS> symks)=0;
virtual void symbol_bind(Item *item, State *state, PgfSymbol sym)=0;
virtual void suspend(Cont *cont, Item *item, bool do_predict, size_t n_suspended,ref<PgfSymbolCat> symcat)=0;
virtual void suspend(Cont *cont, Item *item, bool do_predict, ref<PgfSymbolCat> symcat,interval_t value_i,interval_t lin_idx_i)=0;
virtual void final_item(State *state,CCat *ccat,Item *item,interval_t value,interval_t lin_idx)=0;
virtual void bu_predict(State *state, CCat *ccat)=0;
@@ -247,6 +244,13 @@ protected:
void td_predict(State *state, Cont *cont, Production *prod, Item *xitem, ref<PgfSymbolCat> symcat);
void combine(State *state, Item *item, CCat *ccat);
static
bool instantiate(ref<PgfConcrRule> rule1, size_t *values1, ref<PgfLParam> lparam1,
ref<PgfConcrRule> rule2, size_t *values2, ref<PgfLParam> lparam2);
static
interval_t interval(ref<PgfConcrRule> rule, size_t *values, ref<PgfLParam> lparam);
void get_info(CCat *ccat, ref<PgfConcrRule> *rule, size_t **pvalues);
CCat *get_epsilon_ccat(PgfText *name, PgfMetaId fid);
@@ -272,7 +276,7 @@ class PGF_INTERNAL_DECL PgfParser : private PgfAbstractParser, public PgfExprEnu
virtual State *new_state(const PgfTextSpot &start);
virtual void symbol_token(Item *item, State *state, ref<PgfSymbolKS> symks);
virtual void symbol_bind(Item *item, State *state, PgfSymbol sym);
virtual void suspend(Cont *cont,Item *item,bool do_predict,size_t n_suspended,ref<PgfSymbolCat> symcat);
virtual void suspend(Cont *cont,Item *item,bool do_predict,ref<PgfSymbolCat> symcat,interval_t value_i,interval_t lin_idx_i);
virtual void final_item(State *state,CCat *ccat,Item *item,interval_t value,interval_t lin_idx);
virtual void bu_predict(State *state, CCat *ccat);
@@ -314,7 +318,7 @@ private:
virtual State *new_state(const PgfTextSpot &start);
virtual void symbol_token(Item *item, State *state, ref<PgfSymbolKS> symks);
virtual void symbol_bind(Item *item, State *state, PgfSymbol sym);
virtual void suspend(Cont *cont, Item *item, bool do_predict, size_t n_suspended,ref<PgfSymbolCat> symcat);
virtual void suspend(Cont *cont, Item *item, bool do_predict, ref<PgfSymbolCat> symcat,interval_t value_i,interval_t lin_idx_i);
virtual void final_item(State *state, CCat *ccat,Item *item,interval_t value,interval_t lin_idx);
virtual void bu_predict(State *state, CCat *ccat);