diff --git a/requirements.txt b/requirements.txt index 3ed1fc1..2759b7a 100644 --- a/requirements.txt +++ b/requirements.txt @@ -5,7 +5,6 @@ jdcal==1.3 odfpy==1.3.5 olefile==0.44 openpyxl==2.4.9 -pkg-resources==0.0.0 pytz==2017.3 records==0.5.2 SQLAlchemy==1.1.14 diff --git a/sqlnet/lib/dbengine.py b/sqlnet/lib/dbengine.py index 619c8b1..2acbe27 100644 --- a/sqlnet/lib/dbengine.py +++ b/sqlnet/lib/dbengine.py @@ -18,7 +18,7 @@ def __init__(self, fdb): def execute_query(self, table_id, query, *args, **kwargs): return self.execute(table_id, query.sel_index, query.agg_index, query.conditions, *args, **kwargs) - def execute(self, table_id, select_index, aggregation_index, conditions, lower=True): + def execute(self, table_id, select_index, aggregation_index, conditions, lower=True, isGold=False): if not table_id.startswith('table'): table_id = 'table_{}'.format(table_id.replace('-', '_')) table_info = self.db.query('SELECT sql from sqlite_master WHERE tbl_name = :name', name=table_id).all()[0].sql.replace('\n','') @@ -47,6 +47,12 @@ def execute(self, table_id, select_index, aggregation_index, conditions, lower=T if where_clause: where_str = 'WHERE ' + ' AND '.join(where_clause) query = 'SELECT {} AS result FROM {} {}'.format(select, table_id, where_str) + + if isGold: + print("Gold query: " + query) + else: + print("Pred query: " + query) + #print query out = self.db.query(query, **where_map) return [o.result for o in out] diff --git a/sqlnet/utils.py b/sqlnet/utils.py index 2311671..939de88 100644 --- a/sqlnet/utils.py +++ b/sqlnet/utils.py @@ -127,6 +127,11 @@ def to_batch_query(sql_data, idxes, st, ed): table_ids.append(sql_data[idxes[i]]['table_id']) return query_gt, table_ids +def pretty_print(vis_data): + print('question:', vis_data[0]) + print('headers: (%s)'%(' || '.join(vis_data[1]))) + print('query:', vis_data[2]) + def epoch_train(model, optimizer, batch_size, sql_data, table_data, pred_entry): model.train() perm=np.random.permutation(len(sql_data)) @@ -137,12 +142,13 @@ def epoch_train(model, optimizer, batch_size, sql_data, table_data, pred_entry): q_seq, col_seq, col_num, ans_seq, query_seq, gt_cond_seq = \ to_batch_seq(sql_data, table_data, perm, st, ed) + gt_where_seq = model.generate_gt_where_seq(q_seq, col_seq, query_seq) gt_sel_seq = [x[1] for x in ans_seq] score = model.forward(q_seq, col_seq, col_num, pred_entry, gt_where=gt_where_seq, gt_cond=gt_cond_seq, gt_sel=gt_sel_seq) loss = model.loss(score, ans_seq, pred_entry, gt_where_seq) - cum_loss += loss.data.cpu().numpy()[0]*(ed - st) + cum_loss += loss.data.cpu().numpy()*(ed - st) optimizer.zero_grad() loss.backward() optimizer.step() @@ -175,11 +181,12 @@ def epoch_exec_acc(model, batch_size, sql_data, table_data, db_path): for idx, (sql_gt, sql_pred, tid) in enumerate( zip(query_gt, pred_queries, table_ids)): + print("question:", sql_data[idx]["question"]) ret_gt = engine.execute(tid, - sql_gt['sel'], sql_gt['agg'], sql_gt['conds']) + sql_gt['sel'], sql_gt['agg'], sql_gt['conds'], isGold=True) try: ret_pred = engine.execute(tid, - sql_pred['sel'], sql_pred['agg'], sql_pred['conds']) + sql_pred['sel'], sql_pred['agg'], sql_pred['conds'], isGold=False) except: ret_pred = None tot_acc_num += (ret_gt == ret_pred) @@ -198,6 +205,8 @@ def epoch_acc(model, batch_size, sql_data, table_data, pred_entry): ed = st+batch_size if st+batch_size < len(perm) else len(perm) q_seq, col_seq, col_num, ans_seq, query_seq, gt_cond_seq, raw_data = to_batch_seq(sql_data, table_data, perm, st, ed, ret_vis_data=True) + # print(raw_data[0]) + raw_q_seq = [x[0] for x in raw_data] raw_col_seq = [x[1] for x in raw_data] query_gt, table_ids = to_batch_query(sql_data, perm, st, ed) @@ -242,10 +251,10 @@ def epoch_reinforce_train(model, optimizer, batch_size, sql_data, table_data, db for idx, (sql_gt, sql_pred, tid) in enumerate( zip(query_gt, pred_queries, table_ids)): ret_gt = engine.execute(tid, - sql_gt['sel'], sql_gt['agg'], sql_gt['conds']) + sql_gt['sel'], sql_gt['agg'], sql_gt['conds'], isGold=True) try: ret_pred = engine.execute(tid, - sql_pred['sel'], sql_pred['agg'], sql_pred['conds']) + sql_pred['sel'], sql_pred['agg'], sql_pred['conds'], isGold=False) except: ret_pred = None diff --git a/train.py b/train.py index ed0cab5..64477f9 100644 --- a/train.py +++ b/train.py @@ -82,7 +82,7 @@ print "Init dev acc_ex: %s"%epoch_exec_acc( model, BATCH_SIZE, val_sql_data, val_table_data, DEV_DB) torch.save(model.cond_pred.state_dict(), cond_m) - for i in range(100): + for i in range(1): print 'Epoch %d @ %s'%(i+1, datetime.datetime.now()) print ' Avg reward = %s'%epoch_reinforce_train( model, optimizer, BATCH_SIZE, sql_data, table_data, TRAIN_DB) @@ -121,7 +121,7 @@ torch.save(model.cond_pred.state_dict(), cond_m) if args.train_emb: torch.save(model.cond_embed_layer.state_dict(), cond_e) - for i in range(100): + for i in range(1): print 'Epoch %d @ %s'%(i+1, datetime.datetime.now()) print ' Loss = %s'%epoch_train( model, optimizer, BATCH_SIZE,