implemented A* search

This commit is contained in:
Krasimir Angelov
2026-08-19 11:41:09 +02:00
parent 56026271ff
commit 6c2aeb6b95
2 changed files with 41 additions and 33 deletions
+31 -24
View File
@@ -317,7 +317,7 @@ void PgfAbstractParser::complete(Item *item, State *state)
break; break;
if (ccat->prods.size() == 1) { if (ccat->prods.size() == 1) {
bu_predict(state, ccat); bu_predict(state, item->outside_prob, ccat);
for (auto it1 : ccat->cont->suspended.overlaps(ccat->value)) { for (auto it1 : ccat->cont->suspended.overlaps(ccat->value)) {
for (auto it2 : it1.second.overlaps(ccat->lin_idx)) { for (auto it2 : it1.second.overlaps(ccat->lin_idx)) {
@@ -899,7 +899,7 @@ PgfParser::~PgfParser()
} }
void PgfParser::bu_predict(PgfPhrasetable<PgfSymbolKS> phrasetable, void PgfParser::bu_predict(PgfPhrasetable<PgfSymbolKS> phrasetable,
State *state, State *state, prob_t outside_prob,
ptrdiff_t min, ptrdiff_t max) ptrdiff_t min, ptrdiff_t max)
{ {
if (phrasetable == 0) if (phrasetable == 0)
@@ -908,40 +908,41 @@ void PgfParser::bu_predict(PgfPhrasetable<PgfSymbolKS> phrasetable,
PgfTextSpot current = state->end; PgfTextSpot current = state->end;
int cmp = text_symbol_cmp(&current,end,phrasetable->value.key,case_sensitive); int cmp = text_symbol_cmp(&current,end,phrasetable->value.key,case_sensitive);
if (cmp < 0) { if (cmp < 0) {
bu_predict(phrasetable->left,state,min,max); bu_predict(phrasetable->left,state,outside_prob,min,max);
} else if (cmp > 0) { } else if (cmp > 0) {
ptrdiff_t len = current.ptr - state->end.ptr; ptrdiff_t len = current.ptr - state->end.ptr;
if (min <= len-1) if (min <= len-1)
bu_predict(phrasetable->left,state,min,len-1); bu_predict(phrasetable->left,state,outside_prob,min,len-1);
if (len <= max) if (len <= max)
bu_predict(phrasetable->right,state,len,max); bu_predict(phrasetable->right,state,outside_prob,len,max);
} else { } else {
ptrdiff_t len = current.ptr - state->end.ptr; ptrdiff_t len = current.ptr - state->end.ptr;
if (min <= len) if (min <= len)
bu_predict(phrasetable->left,state,min,len); bu_predict(phrasetable->left,state,outside_prob,min,len);
if (len > 0) { if (len > 0) {
State *next_state = new_state(current);
for (size_t i = 0; i < phrasetable->value.n_items; i++) { for (size_t i = 0; i < phrasetable->value.n_items; i++) {
std::map<ref<PgfConcrLincat>, bool> visited; //std::map<ref<PgfConcrLincat>, bool> visited;
//if (!td_reachable(state, phrasetable->items[i], visited)) //if (!td_reachable(state, phrasetable->items[i], visited))
// continue; // continue;
Item *item = bu_item(state, phrasetable->value.items[i]); Item *item = bu_item(state, outside_prob, phrasetable->value.items[i]);
item->dot++; item->dot++;
State *next_state = new_state(current,item->outside_prob+item->inside_prob);
next_state->push_item(item); next_state->push_item(item);
} }
} }
if (len <= max) if (len <= max)
bu_predict(phrasetable->right,state,len,max); bu_predict(phrasetable->right,state,outside_prob,len,max);
} }
} }
void PgfParser::bu_predict(PgfPhrasetable<PgfSymbolBIND> phrasetable, void PgfParser::bu_predict(PgfPhrasetable<PgfSymbolBIND> phrasetable,
State *state) State *state, prob_t outside_prob)
{ {
size_t n_items = 0; size_t n_items = 0;
vector<ref<PgfItem>> items = vector<ref<PgfItem>> items =
@@ -956,6 +957,7 @@ void PgfParser::bu_predict(PgfPhrasetable<PgfSymbolBIND> phrasetable,
next_state->end = state->end; next_state->end = state->end;
next_state->next = state->next; next_state->next = state->next;
next_state->needs_bind = false; next_state->needs_bind = false;
next_state->viterbi_prob = state->viterbi_prob;
state->next = next_state; state->next = next_state;
} }
@@ -963,13 +965,13 @@ void PgfParser::bu_predict(PgfPhrasetable<PgfSymbolBIND> phrasetable,
//std::map<ref<PgfConcrLincat>, bool> visited; //std::map<ref<PgfConcrLincat>, bool> visited;
//if (!td_reachable(state, phrasetable->items[i], visited)) //if (!td_reachable(state, phrasetable->items[i], visited))
// continue; // continue;
Item *item = bu_item(state, items[i]); Item *item = bu_item(state, outside_prob, items[i]);
item->dot++; item->dot++;
next_state->push_item(item); next_state->push_item(item);
} }
} }
void PgfParser::bu_predict(State *state, CCat *ccat) void PgfParser::bu_predict(State *state, prob_t outside_prob, CCat *ccat)
{ {
size_t n_items = 0; size_t n_items = 0;
vector<ref<PgfItem>> items = 0; vector<ref<PgfItem>> items = 0;
@@ -987,7 +989,7 @@ void PgfParser::bu_predict(State *state, CCat *ccat)
//std::map<ref<PgfConcrLincat>, bool> visited; //std::map<ref<PgfConcrLincat>, bool> visited;
//if (!td_reachable(ccat->cont->state, items[i], visited)) //if (!td_reachable(ccat->cont->state, items[i], visited))
// continue; // continue;
auto new_item = bu_item(ccat->cont->state, items[i]); auto new_item = bu_item(ccat->cont->state, outside_prob, items[i]);
combine(state,new_item,ccat); combine(state,new_item,ccat);
} }
} }
@@ -1023,7 +1025,7 @@ bool PgfParser::td_reachable(State *state, ref<PgfItem> pitem,
return false; return false;
} }
PgfAbstractParser::Item *PgfParser::bu_item(State *state, ref<PgfItem> pitem) PgfAbstractParser::Item *PgfParser::bu_item(State *state, prob_t outside_prob, ref<PgfItem> pitem)
{ {
Item *item = NULL; Item *item = NULL;
@@ -1061,7 +1063,7 @@ PgfAbstractParser::Item *PgfParser::bu_item(State *state, ref<PgfItem> pitem)
item->syms = pitem->rule->syms.as_vector(); item->syms = pitem->rule->syms.as_vector();
item->rule = pitem->rule; item->rule = pitem->rule;
item->inside_prob = lin->absfun->prob; item->inside_prob = lin->absfun->prob;
item->outside_prob = 0; item->outside_prob = outside_prob;
for (size_t i = 0; i < pitem->args.size(); i++) { for (size_t i = 0; i < pitem->args.size(); i++) {
item->args[i] = 0; item->args[i] = 0;
@@ -1149,7 +1151,7 @@ void PgfParser::prepare(ref<PgfConcrLincat> start)
#endif #endif
PgfTextSpot start_spot = {0, (uint8_t *) sentence->text}; PgfTextSpot start_spot = {0, (uint8_t *) sentence->text};
State *state = new_state(start_spot); State *state = new_state(start_spot, 0);
for (size_t i = start->n_lindefs; i < start->rules.size(); i++) { for (size_t i = start->n_lindefs; i < start->rules.size(); i++) {
ref<PgfConcrRule> rule = start->rules[i]; ref<PgfConcrRule> rule = start->rules[i];
@@ -1183,7 +1185,8 @@ PgfExpr PgfParser::fetch(PgfDB *db, prob_t *prob)
while (state != NULL) { while (state != NULL) {
if (state->queue.size() > 0) { if (state->queue.size() > 0) {
Item *item = state->queue.front(); Item *item = state->queue.front();
prob_t prob = item->outside_prob + item->inside_prob; prob_t delta = current_state->viterbi_prob - state->viterbi_prob;
prob_t prob = item->outside_prob + item->inside_prob + delta;
if (min_prob > prob) { if (min_prob > prob) {
min_prob = prob; min_prob = prob;
min_state = state; min_state = state;
@@ -1360,7 +1363,7 @@ PgfExpr PgfParser::process_expr(ExprState *estate, prob_t *prob)
return 0; return 0;
} }
PgfAbstractParser::State *PgfParser::new_state(const PgfTextSpot &start) PgfAbstractParser::State *PgfParser::new_state(const PgfTextSpot &start, prob_t viterbi_prob)
{ {
State **prev = &current_state; State **prev = &current_state;
State *state = current_state; State *state = current_state;
@@ -1374,6 +1377,7 @@ PgfAbstractParser::State *PgfParser::new_state(const PgfTextSpot &start)
state = new State; state = new State;
state->start = start; state->start = start;
state->end = start; state->end = start;
state->viterbi_prob = viterbi_prob;
state->next = *prev; state->next = *prev;
*prev = state; *prev = state;
@@ -1397,7 +1401,7 @@ void PgfParser::symbol_token(Item *item, State *state, ref<PgfSymbolKS> symks)
if (text_symbol_cmp(&next,end,symks,case_sensitive) != 0) if (text_symbol_cmp(&next,end,symks,case_sensitive) != 0)
return; return;
State *next_state = new_state(next); State *next_state = new_state(next, item->inside_prob+item->outside_prob);
item->dot++; item->dot++;
process(item, next_state); process(item, next_state);
@@ -1413,6 +1417,7 @@ void PgfParser::symbol_bind(Item *item, State *state, PgfSymbol sym)
next_state->end = state->end; next_state->end = state->end;
next_state->next = state->next; next_state->next = state->next;
next_state->needs_bind = false; next_state->needs_bind = false;
next_state->viterbi_prob = state->viterbi_prob;
state->next = next_state; state->next = next_state;
} }
item->dot++; item->dot++;
@@ -1472,10 +1477,11 @@ void PgfParser::suspend(Cont *cont,Item *item,bool do_predict,ref<PgfSymbolCat>
} }
if (do_predict) { if (do_predict) {
prob_t viterbi_prob = item->inside_prob+item->outside_prob;
if (cont->state->needs_bind) { if (cont->state->needs_bind) {
bu_predict(concr->phrasetable4, cont->state); bu_predict(concr->phrasetable4, cont->state, viterbi_prob);
} else { } else {
bu_predict(concr->phrasetable1, cont->state, 1, sentence->size); bu_predict(concr->phrasetable1, cont->state, viterbi_prob, 1, sentence->size);
} }
} }
} else { } else {
@@ -1606,6 +1612,7 @@ PgfParseTableMaker::PgfParseTableMaker(ref<PgfConcr> concr)
current_state->start.pos = 0; current_state->start.pos = 0;
current_state->start.ptr = NULL; current_state->start.ptr = NULL;
current_state->end = current_state->start; current_state->end = current_state->start;
current_state->viterbi_prob = 0;
current_state->next = NULL; current_state->next = NULL;
} }
@@ -1629,7 +1636,7 @@ ref<PgfItem> PgfParseTableMaker::clone_item(Item *item)
return pitem; return pitem;
} }
PgfAbstractParser::State *PgfParseTableMaker::new_state(const PgfTextSpot &start) PgfAbstractParser::State *PgfParseTableMaker::new_state(const PgfTextSpot &start, prob_t viterbi_prob)
{ {
return current_state; return current_state;
} }
@@ -1716,7 +1723,7 @@ void PgfParseTableMaker::final_item(State *state, CCat *ccat, Item *item, interv
} }
} }
void PgfParseTableMaker::bu_predict(State *state, CCat *ccat) void PgfParseTableMaker::bu_predict(State *state, prob_t outside_prob, CCat *ccat)
{ {
} }
+10 -9
View File
@@ -103,6 +103,7 @@ protected:
std::map<CCat*,Cont*> conts2; std::map<CCat*,Cont*> conts2;
std::map<Cont*,interval_map<interval_map<CCat*>>> completed; std::map<Cont*,interval_map<interval_map<CCat*>>> completed;
std::vector<Item*> queue; std::vector<Item*> queue;
prob_t viterbi_prob;
State *next; State *next;
@@ -237,12 +238,12 @@ protected:
void symbol(Item *item, State *state, PgfSymbol sym); void symbol(Item *item, State *state, PgfSymbol sym);
void complete(Item *item, State *state); void complete(Item *item, State *state);
virtual State *new_state(const PgfTextSpot &start)=0; virtual State *new_state(const PgfTextSpot &start, prob_t viterbi_prob)=0;
virtual void symbol_token(Item *item, State *state, ref<PgfSymbolKS> symks)=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 symbol_bind(Item *item, State *state, PgfSymbol sym)=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 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 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; virtual void bu_predict(State *state, prob_t outside_prob, CCat *ccat)=0;
void td_epsilon(State *state, Cont *cont, ref<PgfItem> pitem, Item *xitem, ref<PgfSymbolCat> symcat); void td_epsilon(State *state, Cont *cont, ref<PgfItem> pitem, Item *xitem, ref<PgfSymbolCat> symcat);
void td_predict(State *state, Cont *cont, Production *prod, Item *xitem, ref<PgfSymbolCat> symcat); void td_predict(State *state, Cont *cont, Production *prod, Item *xitem, ref<PgfSymbolCat> symcat);
@@ -277,20 +278,20 @@ class PGF_INTERNAL_DECL PgfParser : private PgfAbstractParser, public PgfExprEnu
uint8_t *end; uint8_t *end;
bool case_sensitive; bool case_sensitive;
virtual State *new_state(const PgfTextSpot &start); virtual State *new_state(const PgfTextSpot &start, prob_t viterbi_prob);
virtual void symbol_token(Item *item, State *state, ref<PgfSymbolKS> symks); virtual void symbol_token(Item *item, State *state, ref<PgfSymbolKS> symks);
virtual void symbol_bind(Item *item, State *state, PgfSymbol sym); virtual void symbol_bind(Item *item, State *state, PgfSymbol sym);
virtual void suspend(Cont *cont,Item *item,bool do_predict,ref<PgfSymbolCat> symcat,interval_t value_i,interval_t lin_idx_i); 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 final_item(State *state,CCat *ccat,Item *item,interval_t value,interval_t lin_idx);
virtual void bu_predict(State *state, CCat *ccat); virtual void bu_predict(State *state, prob_t outside_prob, CCat *ccat);
void bu_predict(PgfPhrasetable<PgfSymbolBIND> phrasetable, State *state); void bu_predict(PgfPhrasetable<PgfSymbolBIND> phrasetable, State *state, prob_t outside_prob);
void bu_predict(PgfPhrasetable<PgfSymbolKS> phrasetable, State *state, ptrdiff_t min, ptrdiff_t max); void bu_predict(PgfPhrasetable<PgfSymbolKS> phrasetable, State *state, prob_t outside_prob, ptrdiff_t min, ptrdiff_t max);
void make_chunks(State *state, std::vector<CCat*> &chunks, prob_t prob); void make_chunks(State *state, std::vector<CCat*> &chunks, prob_t prob);
PgfExpr process_expr(ExprState *estate, prob_t *prob); PgfExpr process_expr(ExprState *estate, prob_t *prob);
bool td_reachable(State *state, ref<PgfItem> pitem, std::map<ref<PgfConcrLincat>, bool> &visited); bool td_reachable(State *state, ref<PgfItem> pitem, std::map<ref<PgfConcrLincat>, bool> &visited);
Item *bu_item(State *state, ref<PgfItem> pitem); Item *bu_item(State *state, prob_t outside_prob, ref<PgfItem> pitem);
static static
void print_expr_state_left(PgfPrinter *printer, PgfMarshaller *m, ExprState *estate); void print_expr_state_left(PgfPrinter *printer, PgfMarshaller *m, ExprState *estate);
@@ -319,12 +320,12 @@ public:
class PGF_INTERNAL_DECL PgfParseTableMaker : private PgfAbstractParser class PGF_INTERNAL_DECL PgfParseTableMaker : private PgfAbstractParser
{ {
private: private:
virtual State *new_state(const PgfTextSpot &start); virtual State *new_state(const PgfTextSpot &start, prob_t viterbi_prob);
virtual void symbol_token(Item *item, State *state, ref<PgfSymbolKS> symks); virtual void symbol_token(Item *item, State *state, ref<PgfSymbolKS> symks);
virtual void symbol_bind(Item *item, State *state, PgfSymbol sym); virtual void symbol_bind(Item *item, State *state, PgfSymbol sym);
virtual void suspend(Cont *cont, Item *item, bool do_predict, ref<PgfSymbolCat> symcat,interval_t value_i,interval_t lin_idx_i); 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 final_item(State *state, CCat *ccat,Item *item,interval_t value,interval_t lin_idx);
virtual void bu_predict(State *state, CCat *ccat); virtual void bu_predict(State *state, prob_t outside_prob, CCat *ccat);
static static
ref<PgfItem> clone_item(Item *item); ref<PgfItem> clone_item(Item *item);