Skip to content
This repository was archived by the owner on Nov 4, 2025. It is now read-only.
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
54 changes: 37 additions & 17 deletions geosolver/run.py
Original file line number Diff line number Diff line change
Expand Up @@ -120,18 +120,20 @@ def _annotated_unit_test(query):

ans = solve(reduced_formulas, choice_formulas, assignment=core_parse.variable_assignment)
print "ans:", ans

if choice_formulas is None:
attempted = True
if abs(ans - float(question.answer)) < 0.01:
correct = True
print("No choices correct guess: {0} <golden> : {1}".format(ans, quesion.answer))
else:
correct = False
else:
attempted = True
c_pair = max(ans.iteritems(), key=lambda pair: pair[1].conf)
c = max(ans.iteritems(), key=lambda pair: pair[1].conf)[0]
if c == int(question.answer):
correct = True
print("With choices correct guess: {0} <golden> : {1}, value: {2}".format(c, quesion.answer, c_pair))
else:
correct = False

Expand All @@ -158,6 +160,7 @@ def handler(signum, frame):
try:
result = _full_unit_test(combined_model, question, label_data)
except Exception, e:
print("Exceptiont in full_unit_test")
logging.error(question.key)
logging.exception(e)
result = SimpleResult(question.key, True, False, False)
Expand Down Expand Up @@ -369,23 +372,25 @@ def _full_unit_test(combined_model, question, label_data):
json.dump(entity_list, open(entity_list_path, 'wb'))
json.dump(solution, open(solution_path, 'wb'))

return SimpleResult(question.key, False, False, True) # Early termination
#return SimpleResult(question.key, False, False, True) # Early termination

print "Solving..."
ans = solve(reduced_formulas, choice_formulas, assignment=None)#core_parse.variable_assignment)
print "ans:", ans


if choice_formulas is None:
penalized = False
if Equals(ans, float(question.answer)).conf > 0.98:
print("Correct for non-multiple choice: {0}, golden: {1}".format(ans, quetion.answer))
correct = True
else:
correct = False
else:
idx, tv = max(ans.iteritems(), key=lambda pair: pair[1].conf)
if tv.conf > 0.98:
if idx == int(question.answer):
print("index: {0}, answer_index: {1}".format(idx, question.answer))
if idx == int(float(question.answer)):
print("Correct for multiple choice: index: {0} value: {1}, golden: {2}".format(idx, tv, question.answer))
correct = True
penalized = False
else:
Expand Down Expand Up @@ -483,12 +488,22 @@ def full_test():
te_ids = ids1+ids2+ids3
te_ids = ids4+ids6

load = True
load = False

tr_questions = geoserver_interface.download_questions('aaai')
te_questions = geoserver_interface.download_questions('emnlp')
te_keys = [968, 971, 973, 1018]
all_questions = dict(tr_questions.items() + te_questions.items())
te_questions = geoserver_interface.download_questions('emnlp')
to_questions = geoserver_interface.download_questions('official')

te_keys = []
for x in sys.argv[1].split(","):
if x.isalpha():
temp_q = geoserver_interface.download_questions(x)
te_questions = dict(te_questions.items() + temp_q.items())
te_keys = te_keys + temp_q.keys()
elif x.isdigit():
te_keys.append(int(x))

all_questions = dict(tr_questions.items() + te_questions.items() + to_questions.items())
tr_ids = tr_questions.keys()
te_ids = te_questions.keys()

Expand Down Expand Up @@ -517,16 +532,21 @@ def full_test():
else:
cm = pickle.load(open('cm.p', 'rb'))

print "test ids: %s" % ", ".join(str(k) for k in te_s.keys())
print "Test ids: %s" % ", ".join(str(k) for k in te_s.keys())
for idx, id_ in enumerate(te_keys):
question = all_questions[id_]
label = all_labels[id_]
id_ = str(id_)
print "-"*80
print "id: %s" % id_
result = full_unit_test(cm, question, label)
print result.message
print result
if id_ not in all_labels:
print "No annotation for image"
result = SimpleResult(id_, True, False, False)
else:
question = all_questions[id_]
label = all_labels[id_]
id_ = str(id_)
print "-"*80
print "id: %s" % id_
result = full_unit_test(cm, question, label)
print("RESULT: {0} ".format(result))
print result.message
print result
if result.error:
error += 1
if result.penalized:
Expand Down
3 changes: 2 additions & 1 deletion geosolver/solver/numeric_solver.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
import functools
import algopy
import pyipopt

from scipy.optimize import minimize, newton_krylov, basinhopping
import numpy as np
from scipy.optimize.nonlin import NoConvergence
Expand Down Expand Up @@ -51,6 +51,7 @@ def evaluate(self, variable_node, th=None):
variable_node = self.variable_handler.add(variable_node)
if not self.assigned:
self.solve()
# print('evi', variable_node, self.assignment)
return evaluate(variable_node, self.assignment)


Expand Down
8 changes: 7 additions & 1 deletion geosolver/solver/solve.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@ def solve(given_formulas, choice_formulas=None, assignment=None):
#1. Find query formula in true formulas
true_formulas = []
query_formula = None
print("Given formulas: {0}".format(given_formulas))
for formula in given_formulas:
assert isinstance(formula, FormulaNode)
if formula.has_signature("What") or formula.has_signature("Which") or formula.has_signature("Find"):
Expand All @@ -41,6 +42,8 @@ def solve(given_formulas, choice_formulas=None, assignment=None):
for key, choice_formula in choice_formulas.iteritems():
equal_formula = FormulaNode(signatures['Equals'], [ns.assignment['What'], choice_formula])
out[key] = ns.evaluate(equal_formula)
print equal_formula
print ns.assignment["What"]

"""
ns = NumericSolver(true_formulas)
Expand All @@ -67,6 +70,9 @@ def solve(given_formulas, choice_formulas=None, assignment=None):
for key, choice_formula in choice_formulas.iteritems():
replaced_formula = FormulaNode(signatures['Equals'], [query_formula.children[0], choice_formula])
out[key] = ns.evaluate(replaced_formula)
out[key].formula = replaced_formula
print replaced_formula
print("Key: ", out[key])
# display_entities(ns)


Expand Down Expand Up @@ -99,4 +105,4 @@ def solve(given_formulas, choice_formulas=None, assignment=None):
end_time = time.time()
delta_time = end_time - start_time
print "%.2f seconds" % delta_time
return out
return out
Loading