File size: 8,159 Bytes
b296ad4 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 | """Training-only schema-layout, runtime-name and implicit city-language augmentation."""
import argparse
import copy
import hashlib
import json
from pathlib import Path
import random
import re
import shutil
import sqlite3
import string
import sqlglot
from sqlglot import exp
from tinyquery.data import SQL_OPS,build_record,compact,ddls,gold_sql,load_templates,serialize,sqlite_sql
def layout(row,rng):
a=row['target'].get('arguments',{});query=next((a[k] for k in ['sql','query'] if str(a.get(k,'')).startswith('SELECT ')),None)
required=set();tables=set()
if query:
tree=sqlglot.parse_one(query,read='postgres' if row['backend']=='supabase' else 'mysql')
required={c.name for c in tree.find_all(exp.Column)};tables={t.name for t in tree.find_all(exp.Table)}
schemas=[]
for ddl in row['context']['schema']:
tree=sqlglot.parse_one(ddl);schema=tree.this
if query and schema.this.name not in tables and rng.random()<.8:continue
columns=[]
for column in schema.expressions:
if not isinstance(column,exp.ColumnDef):continue
if column.name not in required and column.name!='id' and rng.random()<.4:continue
column.set('constraints',[c for c in column.args.get('constraints',[]) if not isinstance(c.kind,exp.Reference)])
columns.append(column)
rng.shuffle(columns);schema.set('expressions',columns);schemas.append(tree.sql()+';')
rng.shuffle(schemas);row['context']['schema']=schemas
if query:
db=sqlite3.connect(':memory:');db.create_function('YEAR',1,lambda x:0);db.create_function('MONTH',1,lambda x:0)
try:
for ddl in schemas:db.execute(ddl)
db.execute(sqlite_sql(sqlglot.parse_one(query,read='postgres' if row['backend']=='supabase' else 'mysql')))
finally:db.close()
def rename_tools(row,rng):
prefix=''.join(rng.choices(string.ascii_lowercase,k=rng.randrange(5,14)))+rng.choice(['.','__','_'])
for tool in row['context']['tools']:
old=tool['name'];suffix=old.rsplit('.',1)[-1].rsplit('__',1)[-1]
tool['name']=prefix+suffix
if row['target'].get('name')==old:row['target']['name']=tool['name']
def main():
p=argparse.ArgumentParser();p.add_argument('--source',default='data/tinyquery-v6');p.add_argument('--out',default='data/tinyquery-v7')
args=p.parse_args();source=Path(args.source);out=Path(args.out);out.mkdir(parents=True,exist_ok=True)
for name in ['validation.jsonl','test.jsonl','manual.jsonl','tokenizer.json']:shutil.copy2(source/name,out/name)
rng=random.Random(99707);counts={'layout':0,'runtime_names':0,'city_language':0};seen=set()
with (out/'train.jsonl').open('w') as stream:
with (source/'train.jsonl').open() as src:
for line in src:
stream.write(line);seen.add(hashlib.sha256(json.loads(line)['prompt'].encode()).hexdigest())
def emit(row,kind):
row['id']+='_layout_99707';row['sample_weight']=6
row['provenance']+='; '+kind+' augmentation, training inputs only'
row['response']=compact(row['target']);row['prompt']=serialize(row['context'],row['question'])
fingerprint=hashlib.sha256(row['prompt'].encode()).hexdigest()
if fingerprint in seen:return
seen.add(fingerprint);stream.write(json.dumps(row,ensure_ascii=False)+'\n');counts[kind]+=1
# Use only filtered training rows; never render or inspect held-out language.
with (source/'train.jsonl').open() as src:
for line in src:
row=json.loads(line);op=row['operation']
if op in SQL_OPS and rng.random()<.035:
layout(row,rng)
if rng.random()<.5:rename_tools(row,rng)
emit(row,'layout')
elif op not in SQL_OPS and row['target']['action']=='call' and rng.random()<.22:
rename_tools(row,rng);emit(row,'runtime_names')
templates=load_templates('data/tinyquery/templates.jsonl')
domains=['members','attendees','officers','candidates','speakers','volunteers','staff','participants','residents','contacts','visitors','players','customers']
hindi_tables=dict(zip(domains,['सदस्यों','उपस्थित लोगों','अधिकारियों','उम्मीदवारों','वक्ताओं','स्वयंसेवकों','कर्मचारियों','प्रतिभागियों','निवासियों','संपर्कों','आगंतुकों','खिलाड़ियों','ग्राहकों']))
hindi_cities={'Delhi':'दिल्ली','Mumbai':'मुंबई','Pune':'पुणे','Lucknow':'लखनऊ','Jaipur':'जयपुर','Agra':'आगरा','Chennai':'चेन्नई'}
for i in range(1200):
domain=rng.choice(domains);backend=rng.choice(['mysql','supabase']);op=rng.choice(['eq','project_eq','count_eq'])
records=build_record(op,domain,i,backend,'train',templates,rng)
for row in records:
# build_record shares context/slots across its four language variants.
row=copy.deepcopy(row)
s=row['slots'];s['category']='city';s['column']='name';s['value']=rng.choice(['Delhi','Mumbai','Pune','Lucknow','Jaipur','Agra','Chennai'])
row['scenario_id']='city_'+row['scenario_id'];row['id']='city_'+row['id'];row['template_index']=-1
row['context']['schema']=ddls(s)
args_=row['target']['arguments'];key=next(k for k in ['sql','query'] if k in args_);args_[key]=gold_sql(op,s,backend)
choices={
'eq':{'en':['Get the {table} from {value}.','List {table} located in {value}.'],
'noisy_en':['need {value} {table} all details','{table} from {value} pls'],
'hi':['{value} के {table} का पूरा विवरण लाओ।','{table} में {value} वाले सभी रिकॉर्ड चाहिए।'],
'hinglish':['{value} wale {table} ki puri details do.','{table} jo {value} mein hain unko dikhao.']},
'project_eq':{'en':['List names of {table} located in {value}.','For {table} from {value}, give their names.'],
'noisy_en':['{value} {table} names pls','names only {table} from {value}'],
'hi':['{value} वाले {table} के नाम बताइए।','{table} में {value} के लोगों के नाम चाहिए।'],
'hinglish':['{value} wale {table} ke naam do.','{table} jo {value} se hain unke names chahiye.']},
'count_eq':{'en':['How many {table} are from {value}?','Count the {table} located in {value}.'],
'noisy_en':['how many {value} {table} total','count {table} from {value} pls'],
'hi':['{value} के कुल कितने {table} हैं?','{table} में {value} वालों की संख्या बताइए।'],
'hinglish':['{value} wale {table} kitne hain?','{table} mein {value} ka total count batao.']}}
row['question']=rng.choice(choices[op][row['language']]).format(**s)
row['provenance']='programmatic semantics + programmatically authored city-language phrasing'
if row['language']=='hi' and rng.random()<.7:
row['question']=row['question'].replace(s['value'],hindi_cities[s['value']]).replace(domain,hindi_tables[domain])
row['provenance']+='; authored Hindi city/table lexical mapping, learned from examples'
layout(row,rng)
if rng.random()<.5:rename_tools(row,rng)
emit(row,'city_language')
(out/'grounding-stats.json').write_text(json.dumps(counts,indent=2));print(counts)
if __name__=='__main__':main()
|