import logging import re import time import uuid from data_gen import Dataset, RawExample from train_target import EMBED, WITH_TEXT_ONLY import subprocess import json import sys from data import wavtoolm2midi, wavtoolmm2midi import random import concurrent.futures import itertools import psycopg class DatasetFromExternalTool(Dataset): def __init__( self, tool_cmdline, split_num_examples, tool_parallelism, high_watermark=1000 ): super().__init__() self.tool_cmdline = tool_cmdline self.split_num_examples = {int(k): v for k, v in split_num_examples.items()} self.tool_parallelism = tool_parallelism self.high_watermark = high_watermark @classmethod def from_config(cls, common_config, my_config): return cls( my_config["tool_cmdline"], my_config["split_num_examples"], my_config["tool_parallelism"], ) def num_examples(self): return dict(self.split_num_examples) def shuffle(self): pass def stream_examples_impl(self, split, ranksize, order): example_q = [] example_futs = [] rank, size = ranksize with concurrent.futures.ThreadPoolExecutor( max_workers=self.tool_parallelism ) as executor: for example_idx in ( itertools.count() if split == 0 else range(self.split_num_examples[split]) ): while True: # harvest futures new_example_futs = [] # if empty, we must wait for at least one future to complete if len(example_q) == 0 and len(example_futs) > 0: concurrent.futures.wait( example_futs, return_when=concurrent.futures.FIRST_COMPLETED ) for fut in example_futs: if fut.done(): example_q.extend(fut.result()) else: new_example_futs.append(fut) example_futs = new_example_futs # sow futures while ( len(example_q) < self.high_watermark and len(example_futs) < self.tool_parallelism ): example_futs.append(executor.submit(self._get_more_examples)) if len(example_q) > 0: break example = example_q.pop(0) if example_idx % size != rank: continue yield RawExample( id=example_idx, descs=[example["prompt"]], example=self.join_with_anchor( wavtoolmm2midi(example["mainContext"]), wavtoolmm2midi(example["mainContinuation"]), ), accomps=[ wavtoolmm2midi(accomp) for accomp in random.sample( example["accompaniments"], random.randint(0, min(1, len(example["accompaniments"]))), ) ], ) def join_with_anchor(self, a, b): xs = a.copy() ys = b.copy() if len(ys) > 0: ys[0] = dict(ys[0]) ys[0]["descriptionAnchor"] = True xs.extend(ys) return xs def _get_more_examples(self): p = subprocess.Popen( self.tool_cmdline, stdout=subprocess.PIPE, shell=True, ) try: while p.poll() is not None: line = p.stdout.readline() if not line: continue return json.loads(line) line = p.stdout.readline() if not line: return [] return json.loads(line) except json.JSONDecodeError: print( f"Warning: failed to decode tool output JSON: {line}", file=sys.stderr, ) return [] def _get_randomizer(cur, index_name): cur.execute( "SELECT indexdef FROM pg_indexes WHERE indexname = %s", (index_name,), ) indexdef = cur.fetchone() if indexdef is None: raise RuntimeError("index not found") mat = re.match(r".*hashint4extended\(.*,(.*)\)\)$", indexdef[0]) if mat is None: raise RuntimeError("index has wrong definition") return mat.group(1) def _shuffle_randomizer(cur, table_name, index_name): cur.execute(f"DROP INDEX IF EXISTS {index_name}") seed = random.randint(-(2**63), 2**63 - 1) cur.execute( f"CREATE INDEX {index_name} ON {table_name} (hashint4extended(id, {seed}))" ) cur.execute(f"ANALYZE {table_name}") class DatasetFromPostgresExtractedClips(Dataset): def __init__( self, conn_str, eval_split_idx, train_target, min_len, max_len, aug_limit, permit_tags=None, # do not restrict by tag TODO set in config! permit_datasets=None, # all datasets are permitted drum_tags=None, # tags that indicate drum tracks (None means no drums emitted) ): super().__init__() self.conn_str = conn_str self.eval_split_idx = eval_split_idx self.train_target = train_target self.min_len = min_len self.max_len = max_len self.aug_limit = aug_limit self.permit_tags = permit_tags self.permit_datasets = permit_datasets self.drum_tags = drum_tags @classmethod def from_config(cls, common_config, my_config): return cls( my_config["conn_str"], int(common_config["eval_split_idx"]), common_config["train_target"], int(common_config["seq_len_min"]), int(common_config["seq_len_max"]), int(my_config["aug_limit"]), my_config.get("permit_tags", None), my_config.get("permit_datasets", None), my_config.get("drum_tags", None), ) def num_examples(self): with psycopg.connect(self.conn_str) as conn: with conn.cursor() as cur: cur.execute( """ SELECT f.split, count(1) FROM extracted_clips c JOIN files f ON c.file_id = f.id WHERE c.symbolic_length IS NOT NULL AND c.symbolic_length >= %s AND c.symbolic_length <= %s AND ( f.split != %s OR NOT EXISTS ( SELECT FROM automatic_lane_tags lt JOIN tag_values tv ON lt.tag_value_id = tv.id WHERE lt.lane_id = c.lane_id AND tv.value = 'Contaminated Eval' ) ) """ + ( """ AND EXISTS ( SELECT FROM file_descriptions fd WHERE fd.file_id = f.id AND fd.description IS NOT NULL AND fd.description <> '' ) """ if self.train_target == WITH_TEXT_ONLY else "" ) + ( " AND f.dataset_name = ANY(%s)" if self.permit_datasets is not None else "" ) + ( """ AND EXISTS ( SELECT FROM automatic_extracted_clip_tags ct JOIN tag_values tv ON ct.tag_value_id = tv.id WHERE tv.value = ANY(%s) AND ct.extracted_clip_id = c.id ) """ if self.permit_tags is not None else "" ) + " GROUP BY f.split", ( self.min_len, self.max_len, self.eval_split_idx, *( (self.permit_datasets,) if self.permit_datasets is not None else tuple() ), *( (self.permit_tags,) if self.permit_tags is not None else tuple() ), ), ) return {row[0]: row[1] for row in cur.fetchall()} def shuffle(self): with psycopg.connect(self.conn_str) as conn: conn.autocommit = True with conn.transaction(), conn.cursor() as cur: _shuffle_randomizer(cur, "extracted_clips", "extracted_clips_rand_idx") def _emit_drums(self, notes, tags): if any(tag in self.drum_tags for tag in tags): return [dict(n, note=n["note"] + 1000) for n in notes] return notes def _stream_rows(self, split, ranksize, order): with psycopg.connect(self.conn_str) as conn: with conn.cursor() as cur: cur.execute("SET cursor_tuple_fraction = 0.0001") # dank randomizer = _get_randomizer(cur, "extracted_clips_rand_idx") with conn.cursor(name=f"fetch_{uuid.uuid4()}") as cur: start_time = time.time() cur.itersize = 100 cur.execute( """ SELECT c.id, coalesce(fd.descs, '[]'::jsonb) descs, cn.notes example, coalesce(et.tags, '[]'::jsonb) example_all_tags, coalesce(ca.notess, '[]'::jsonb) accomp, coalesce(ca.tagss, '[]'::jsonb) accomp_all_tagss, c.file_id, c.start FROM extracted_clips c JOIN public.files f ON c.file_id = f.id JOIN extracted_clip_notes cn ON cn.extracted_clip_id = c.id LEFT JOIN LATERAL ( SELECT jsonb_agg( jsonb_build_object( 'description', description, 'type', type ) ORDER BY description ) descs FROM file_descriptions fds WHERE fds.file_id = f.id AND description IS NOT NULL AND description <> '' ) fd ON true LEFT JOIN LATERAL ( WITH all_tag_ids AS ( SELECT lt.tag_value_id FROM lane_tags lt WHERE lt.lane_id = c.lane_id UNION SELECT ct.tag_value_id FROM automatic_extracted_clip_tags ct WHERE ct.extracted_clip_id = c.id ) SELECT jsonb_agg(tv.value) tags FROM all_tag_ids atid JOIN tag_values tv ON atid.tag_value_id = tv.id ) et ON true LEFT JOIN LATERAL ( WITH cas AS ( SELECT can.notes, coalesce(cat.tags, '[]'::jsonb) tags FROM extracted_clips ca LEFT JOIN LATERAL ( WITH all_tag_ids AS ( SELECT lt.tag_value_id FROM lane_tags lt WHERE lt.lane_id = ca.lane_id UNION SELECT ct.tag_value_id FROM automatic_extracted_clip_tags ct WHERE ct.extracted_clip_id = ca.id ) SELECT jsonb_agg(tv.value) tags FROM all_tag_ids atid JOIN tag_values tv ON atid.tag_value_id = tv.id ) cat ON true JOIN extracted_clip_notes can ON can.extracted_clip_id = ca.id WHERE ca.file_id = c.file_id AND ca.start = c.start AND ca.lane_id <> c.lane_id AND ca.symbolic_length >= 10 AND c.symbolic_length + ca.symbolic_length <= %s """ + ( """ AND NOT EXISTS ( SELECT FROM automatic_lane_tags lt JOIN tag_values tv ON lt.tag_value_id = tv.id WHERE lt.lane_id = ca.lane_id AND tv.value = 'Contaminated Eval' ) """ if split == self.eval_split_idx else "" ) + " ORDER BY " + ("random()" if order == "random" else "ca.stripe") + """ LIMIT %s ) SELECT jsonb_agg(notes) notess, jsonb_agg(tags) tagss FROM cas ) ca ON true WHERE f.split = %s AND c.stripe %% %s = %s """ + ( "AND fd.descs IS NOT NULL" if self.train_target == WITH_TEXT_ONLY else "" ) + ( " AND f.dataset_name = ANY(%s)" if self.permit_datasets is not None else "" ) + ( """ AND EXISTS ( SELECT FROM automatic_extracted_clip_tags ct JOIN tag_values tv ON ct.tag_value_id = tv.id WHERE tv.value = ANY(%s) AND ct.extracted_clip_id = c.id ) """ if self.permit_tags is not None else "" ) + ( """ AND NOT EXISTS ( SELECT FROM automatic_lane_tags lt JOIN tag_values tv ON lt.tag_value_id = tv.id WHERE lt.lane_id = c.lane_id AND tv.value = 'Contaminated Eval' ) """ if split == self.eval_split_idx else "" ) + """ AND c.symbolic_length >= %s AND c.symbolic_length <= %s ORDER BY """ + ( f"hashint4extended(c.id, {randomizer})" if order == "random" else "c.stripe" ) + " ASC", ( self.max_len, self.aug_limit, split, ranksize[1], ranksize[0], *( (self.permit_datasets,) if self.permit_datasets is not None else tuple() ), *( (self.permit_tags,) if self.permit_tags is not None else tuple() ), self.min_len, self.max_len, ), ) for row in cur: if start_time is not None: logging.info( f"split {split} time to first row {time.time() - start_time}s" ) start_time = None yield row def stream_examples_impl(self, split, ranksize, order): for row in self._stream_rows(split, ranksize, order): yield RawExample( id=row[0], descs=row[1], example=self._emit_drums(wavtoolm2midi(row[2]), row[3]), example_tags=row[3], accomps=[ self._emit_drums(wavtoolm2midi(accomp), accomp_tags) for accomp, accomp_tags in zip(row[4], row[5]) ], accomp_tagss=row[5], ) class DatasetFromPostgresExtractedClipsOverTime(DatasetFromPostgresExtractedClips): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) self._num_examples = None def num_examples(self): if self._num_examples is not None: return self._num_examples with psycopg.connect(self.conn_str) as conn: with conn.cursor() as cur: cur.execute( """ SELECT f.split, count(distinct (c.file_id, c.start)) FROM extracted_clips c JOIN files f ON c.file_id = f.id WHERE c.symbolic_length IS NOT NULL AND c.symbolic_length >= %s AND c.symbolic_length <= %s AND ( f.split != %s OR NOT EXISTS ( SELECT FROM automatic_lane_tags lt JOIN tag_values tv ON lt.tag_value_id = tv.id WHERE lt.lane_id = c.lane_id AND tv.value = 'Contaminated Eval' ) ) """ + ( """ AND EXISTS ( SELECT FROM file_descriptions fd WHERE fd.file_id = f.id AND fd.description IS NOT NULL AND fd.description <> '' ) """ if self.train_target == WITH_TEXT_ONLY else "" ) + ( " AND f.dataset_name = ANY(%s)" if self.permit_datasets is not None else "" ) + ( """ AND EXISTS ( SELECT FROM automatic_extracted_clip_tags ct JOIN tag_values tv ON ct.tag_value_id = tv.id WHERE tv.value = ANY(%s) AND ct.extracted_clip_id = c.id ) """ if self.permit_tags is not None else "" ) + " GROUP BY f.split", ( self.min_len, self.max_len, self.eval_split_idx, *( (self.permit_datasets,) if self.permit_datasets is not None else tuple() ), *( (self.permit_tags,) if self.permit_tags is not None else tuple() ), ), ) self._num_examples = {row[0]: row[1] for row in cur.fetchall()} return self._num_examples def stream_examples_impl(self, split, ranksize, order): seen_file_starts = set() for row in self._stream_rows(split, ranksize, order): # only emit the first example for each (file, start) pair file_start = (row[6], row[7]) if file_start in seen_file_starts: continue seen_file_starts.add(file_start) yield RawExample( id=row[0], descs=row[1], example=self._emit_drums(wavtoolm2midi(row[2]), row[3]), example_tags=row[3], accomps=[ self._emit_drums(wavtoolm2midi(accomp), accomp_tags) for accomp, accomp_tags in zip(row[4], row[5]) ], accomp_tagss=row[5], ) class DatasetFromPostgresLanes(Dataset): def __init__( self, conn_str, eval_split_idx, permit_datasets=None, # all datasets are permitted ): super().__init__() self.conn_str = conn_str self.eval_split_idx = eval_split_idx self.permit_datasets = permit_datasets @classmethod def from_config(cls, common_config, my_config): return cls( my_config["conn_str"], common_config["eval_split_idx"], my_config.get("permit_datasets", None), ) def num_examples(self): with psycopg.connect(self.conn_str) as conn: with conn.cursor() as cur: cur.execute( """ SELECT f.split, count(1) FROM lanes l JOIN instruments i ON i.id = l.instrument_id JOIN files f ON i.file_id = f.id WHERE (f.split != %s OR NOT EXISTS ( SELECT FROM automatic_lane_tags lt JOIN tag_values tv ON lt.tag_value_id = tv.id WHERE lt.lane_id = l.id AND tv.value = 'Contaminated Eval' )) """ + ( " AND f.dataset_name = ANY(%s)" if self.permit_datasets is not None else "" ) + " GROUP BY f.split", ( self.eval_split_idx, *( (self.permit_datasets,) if self.permit_datasets is not None else tuple() ), ), ) return {row[0]: row[1] for row in cur.fetchall()} def shuffle(self): with psycopg.connect(self.conn_str) as conn: conn.autocommit = True with conn.transaction(), conn.cursor() as cur: _shuffle_randomizer(cur, "lanes", "lanes_rand_idx") def stream_examples_impl(self, split, ranksize, order): with psycopg.connect(self.conn_str) as conn: with conn.cursor() as cur: cur.execute("SET cursor_tuple_fraction = 0.0001") # dank randomizer = _get_randomizer(cur, "lanes_rand_idx") with conn.cursor(name=f"fetch_{uuid.uuid4()}") as cur: start_time = time.time() cur.itersize = 100 cur.execute( """ SELECT l.id, a.notes FROM lanes l JOIN LATERAL ( SELECT jsonb_agg(jsonb_build_object('pitch', n.pitch, 'start', n.start, 'end', n."end", 'velocity', n.velocity)) AS notes FROM notes n WHERE n.lane_id = l.id ) a ON true JOIN instruments i ON i.id = l.instrument_id JOIN files f ON i.file_id = f.id """ + ( "WHERE f.split = %s AND l.stripe %% %s = %s " if split is not None else "WHERE l.stripe %% %s = %s " ) + ( """ AND NOT EXISTS ( SELECT FROM automatic_lane_tags lt JOIN tag_values tv ON lt.tag_value_id = tv.id WHERE lt.lane_id = l.id AND tv.value = 'Contaminated Eval' ) """ if split == self.eval_split_idx else "" ) + ( " AND f.dataset_name = ANY(%s)" if self.permit_datasets is not None else "" ) + " ORDER BY " + ( f"hashint4extended(l.id, {randomizer})" if order == "random" else "l.stripe" ) + " ASC", ( *((split,) if split is not None else tuple()), ranksize[1], ranksize[0], *( (self.permit_datasets,) if self.permit_datasets is not None else tuple() ), ), ) for row in cur: if start_time is not None: logging.info( f"split {split} time to first row {time.time() - start_time}s" ) start_time = None yield RawExample( id=row[0], descs=[], example=wavtoolm2midi(row[1]), accomps=[], ) # d = DatasetFromPostgres( # "host=localhost dbname=composer_new_dataset_v3", WITH_TEXT_ONLY, 30, 300, 1 # ) # print(d.num_examples()) # for e, i in zip(d.stream_examples(0), range(10)): # pprint(e)