| """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 |
| |
| 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: |
| |
| 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() |
|
|