Download dune_extraction.py from Respair/dune_codec: direct link, hf CLI and curl.
- Browser
- Download file 54.6 kB
-
https://huggingface.co/Respair/dune_codec/resolve/main/dune_extraction.py
- Command line
-
hf download hf://Respair/dune_codec/dune_extraction.py
-
curl -L -o dune_extraction.py https://huggingface.co/Respair/dune_codec/resolve/main/dune_extraction.py
54.6 kB
| import io | |
| import json | |
| import math | |
| import os | |
| import shutil | |
| import sys | |
| import time | |
| import warnings | |
| from bisect import bisect_left | |
| from concurrent.futures import ThreadPoolExecutor | |
| from pathlib import Path | |
| from codec.audio_processing.dune_codec import load_dune_audio_tokenizer | |
| import librosa | |
| import numpy as np | |
| import pyarrow as pa | |
| import pyarrow.compute as pc | |
| import torch | |
| import torch.nn.functional as F | |
| from datasets import Audio, Dataset, DatasetDict, Features, Sequence, Value, load_dataset, load_from_disk | |
| from tqdm.auto import tqdm | |
| warnings.filterwarnings("ignore") | |
| MODEL_ID = "Respair/dune_codec" | |
| # contains the metadata rows to enrich and a unique bridge key. | |
| DATASET_SOURCE = "/home/ubuntu/data" | |
| DATASET_CONFIG = None | |
| DATASET_SPLIT = "train" | |
| # if your metadata and audio dataset are the same, copy DATASET_SOURCE here, otherwise ensure you have a index column (like `key`) so these two can be joined | |
| AUDIO_DATASET_SOURCE = ["/home/ubuntu/data"] | |
| AUDIO_DATASET_CONFIG = None | |
| AUDIO_DATASET_SPLIT = "train" | |
| KEY_COLUMN = "key" | |
| AUDIO_COLUMN = "audio" | |
| # output columns | |
| DURATION_COLUMN = "duration" | |
| PREQUANT_COLUMN = "latents" | |
| DISCRETE_TOKENS_COLUMN = "codes" | |
| ERROR_COLUMN = "codec_error" | |
| CHECKPOINT_DIR = "/home/ubuntu/out" | |
| CHECKPOINT_INTERVAL_STEPS = 1024 * 1 | |
| HF_MAX_SHARD_SIZE = "2GB" | |
| # What to save? discrete speech tokens, FSQ pre-quant latents or both? | |
| ## for darya or any other flow matching models, FSQ pre-quant latents is enough | |
| SAVE_PREQUANT = True | |
| SAVE_DISCRETE = True | |
| TEST_RUN = 0 # you probably won't need it | |
| BATCH_SIZE = 64 | |
| NUM_WORKERS = 32 | |
| TARGET_SAMPLE_RATE = 22050 | |
| INFERENCE_DTYPE = torch.bfloat16 | |
| MAX_AUDIO_DURATION_SECONDS = 30.0 # if your clips are capped at another length, tweak it here. | |
| # you can ignore these | |
| MIN_DURATION_BUCKETS = 10 | |
| MAX_DURATION_BUCKETS = 15 | |
| DURATION_HISTOGRAM_RESOLUTION_SECONDS = 0.05 | |
| DURATION_BUCKET_ELBOW_TOLERANCE = 0.03 | |
| ARROW_DURATION_SCAN_BATCH_SIZE = 250_000 | |
| DURATION_BUCKET_INDEX_SCAN_BATCH_SIZE = 250_000 | |
| AUDIO_INDEX_BATCH_SIZE = 100_000 | |
| EXPECTED_PREQUANT_DIM = 52 | |
| def load_dataset_source(source, config, split): | |
| source_path = Path(source).expanduser() | |
| if source_path.exists(): | |
| dataset = load_from_disk(str(source_path)) | |
| if isinstance(dataset, DatasetDict): | |
| if split not in dataset: | |
| raise KeyError( | |
| f"Split '{split}' is not present in local dataset " | |
| f"'{source_path}'. Available splits: {list(dataset.keys())}" | |
| ) | |
| dataset = dataset[split] | |
| return dataset | |
| return load_dataset( | |
| source, | |
| config, | |
| split=split, | |
| streaming=False, | |
| ) | |
| def load_input_dataset(): | |
| dataset = load_dataset_source( | |
| source=DATASET_SOURCE, | |
| config=DATASET_CONFIG, | |
| split=DATASET_SPLIT, | |
| ) | |
| if KEY_COLUMN not in dataset.column_names: | |
| raise KeyError( | |
| f"Primary dataset does not contain the bridge column " | |
| f"'{KEY_COLUMN}'. Available columns: {dataset.column_names}" | |
| ) | |
| if AUDIO_COLUMN in dataset.column_names: | |
| dataset = dataset.remove_columns([AUDIO_COLUMN]) | |
| if TEST_RUN > 0: | |
| dataset = dataset.select(range(min(TEST_RUN, len(dataset)))) | |
| return dataset | |
| def load_audio_datasets(): | |
| sources = ( | |
| list(AUDIO_DATASET_SOURCE) | |
| if isinstance(AUDIO_DATASET_SOURCE, (list, tuple)) | |
| else [AUDIO_DATASET_SOURCE] | |
| ) | |
| audio_datasets = [] | |
| for source in sources: | |
| dataset = load_dataset_source( | |
| source=source, | |
| config=AUDIO_DATASET_CONFIG, | |
| split=AUDIO_DATASET_SPLIT, | |
| ) | |
| required_columns = {KEY_COLUMN, AUDIO_COLUMN} | |
| missing_columns = required_columns.difference(dataset.column_names) | |
| if missing_columns: | |
| raise KeyError( | |
| f"Audio dataset '{source}' is missing required columns: " | |
| f"{sorted(missing_columns)}. " | |
| f"Available columns: {dataset.column_names}" | |
| ) | |
| audio_datasets.append( | |
| dataset.cast_column( | |
| AUDIO_COLUMN, | |
| Audio(sampling_rate=TARGET_SAMPLE_RATE), | |
| ) | |
| ) | |
| return audio_datasets | |
| def normalize_bridge_key(key, dataset_role): | |
| if key is None: | |
| raise ValueError( | |
| f"{dataset_role} row has a null '{KEY_COLUMN}' value." | |
| ) | |
| if isinstance(key, str): | |
| return key | |
| if isinstance(key, (int, np.integer)) and not isinstance(key, (bool, np.bool_)): | |
| return str(int(key)) | |
| raise TypeError( | |
| f"{dataset_role} '{KEY_COLUMN}' value must be a string or integer, " | |
| f"received {type(key).__name__}: {key!r}" | |
| ) | |
| def build_audio_key_index(audio_datasets): | |
| key_to_position = {} | |
| for dataset_position, audio_dataset in enumerate(audio_datasets): | |
| key_dataset = audio_dataset.select_columns([KEY_COLUMN]) | |
| number_of_batches = math.ceil( | |
| len(key_dataset) / AUDIO_INDEX_BATCH_SIZE | |
| ) | |
| progress = tqdm( | |
| key_dataset.iter(batch_size=AUDIO_INDEX_BATCH_SIZE), | |
| total=number_of_batches, | |
| desc=f"Indexing audio keys [{dataset_position + 1}/{len(audio_datasets)}]", | |
| unit="batch", | |
| dynamic_ncols=True, | |
| ) | |
| audio_position = 0 | |
| for key_batch in progress: | |
| for raw_key in key_batch[KEY_COLUMN]: | |
| key = normalize_bridge_key(raw_key, "Audio dataset") | |
| if key in key_to_position: | |
| previous_dataset_position, previous_audio_position = ( | |
| key_to_position[key] | |
| ) | |
| raise ValueError( | |
| f"Duplicate normalized audio key {key!r}: found in " | |
| f"audio dataset {previous_dataset_position} at row " | |
| f"{previous_audio_position} and audio dataset " | |
| f"{dataset_position} at row {audio_position}." | |
| ) | |
| key_to_position[key] = (dataset_position, audio_position) | |
| audio_position += 1 | |
| return key_to_position | |
| # audio i/o and batchh | |
| def decode_audio(audio_value, target_sample_rate): | |
| if isinstance(audio_value, dict) and audio_value.get("array") is not None: | |
| wav = np.asarray(audio_value["array"], dtype=np.float32) | |
| sample_rate = int(audio_value["sampling_rate"]) | |
| elif hasattr(audio_value, "get_all_samples"): | |
| samples = audio_value.get_all_samples() | |
| wav = samples.data.cpu().numpy().astype(np.float32, copy=False) | |
| sample_rate = int(samples.sample_rate) | |
| else: | |
| source = audio_value | |
| if isinstance(audio_value, dict): | |
| if audio_value.get("bytes") is not None: | |
| source = io.BytesIO(audio_value["bytes"]) | |
| else: | |
| source = audio_value.get("path") | |
| elif isinstance(audio_value, (bytes, bytearray)): | |
| source = io.BytesIO(audio_value) | |
| if source is None: | |
| raise ValueError("Audio value has no array, bytes, or path.") | |
| wav, sample_rate = librosa.load( | |
| source, | |
| sr=None, | |
| mono=True, | |
| ) | |
| wav = np.asarray(wav, dtype=np.float32) | |
| if wav.ndim > 1: | |
| wav = wav.mean(axis=0) | |
| if sample_rate != target_sample_rate: | |
| wav = librosa.resample( | |
| wav, | |
| orig_sr=sample_rate, | |
| target_sr=target_sample_rate, | |
| ).astype(np.float32, copy=False) | |
| if len(wav) == 0: | |
| raise ValueError("Decoded audio is empty.") | |
| return wav | |
| def prepare_audio_item(item): | |
| position, row, audio_datasets, audio_key_index, sample_rate = item | |
| row = dict(row) | |
| raw_key = row.get(KEY_COLUMN) | |
| try: | |
| key = normalize_bridge_key(raw_key, "Primary") | |
| if key not in audio_key_index: | |
| raise KeyError( | |
| f"No matching audio row was found for key {key!r}." | |
| ) | |
| dataset_position, audio_position = audio_key_index[key] | |
| audio_value = audio_datasets[dataset_position][audio_position][AUDIO_COLUMN] | |
| wav = decode_audio(audio_value, sample_rate) | |
| return position, row, torch.from_numpy(wav).float(), None | |
| except Exception as exc: | |
| return ( | |
| position, | |
| row, | |
| None, | |
| f"ERROR: audio lookup or decoding failed - {exc}", | |
| ) | |
| def get_duration_bucket_index(num_samples, bucket_boundaries_samples): | |
| bucket_index = bisect_left( | |
| bucket_boundaries_samples, | |
| num_samples, | |
| ) | |
| if bucket_index == len(bucket_boundaries_samples): | |
| return None | |
| return bucket_index | |
| def build_metadata_duration_histogram( | |
| metadata_dataset, | |
| sample_rate, | |
| histogram_bin_samples, | |
| ): | |
| if DURATION_COLUMN not in metadata_dataset.column_names: | |
| raise KeyError( | |
| f"Primary dataset does not contain duration column " | |
| f"'{DURATION_COLUMN}'. Available columns: " | |
| f"{metadata_dataset.column_names}" | |
| ) | |
| total_rows = len(metadata_dataset) | |
| if total_rows == 0: | |
| raise ValueError( | |
| "Cannot calculate duration buckets from an empty dataset." | |
| ) | |
| duration_table = metadata_dataset.select_columns( | |
| [DURATION_COLUMN] | |
| ).data | |
| duration_column = duration_table.column(DURATION_COLUMN) | |
| if not isinstance( | |
| duration_column, | |
| (pa.Array, pa.ChunkedArray), | |
| ): | |
| raise TypeError( | |
| f"Expected an Arrow Array or ChunkedArray for " | |
| f"'{DURATION_COLUMN}', received {type(duration_column)}." | |
| ) | |
| maximum_histogram_samples = int( | |
| math.ceil(MAX_AUDIO_DURATION_SECONDS * sample_rate) | |
| ) | |
| histogram_size = ( | |
| (maximum_histogram_samples - 1) // histogram_bin_samples | |
| ) + 1 | |
| counts = np.zeros(histogram_size, dtype=np.int64) | |
| sample_sums = np.zeros(histogram_size, dtype=np.float64) | |
| bin_maxima = np.zeros(histogram_size, dtype=np.int64) | |
| observed_max_samples = 0 | |
| processed_rows = 0 | |
| progress = tqdm( | |
| total=total_rows, | |
| desc="Scanning Arrow durations", | |
| unit="rows", | |
| dynamic_ncols=True, | |
| ) | |
| for batch_start in range( | |
| 0, | |
| total_rows, | |
| ARROW_DURATION_SCAN_BATCH_SIZE, | |
| ): | |
| batch_length = min( | |
| ARROW_DURATION_SCAN_BATCH_SIZE, | |
| total_rows - batch_start, | |
| ) | |
| arrow_batch = duration_column.slice( | |
| batch_start, | |
| batch_length, | |
| ) | |
| if isinstance(arrow_batch, pa.ChunkedArray): | |
| arrow_batch = arrow_batch.combine_chunks() | |
| if arrow_batch.null_count: | |
| null_mask = pc.is_null(arrow_batch).to_numpy( | |
| zero_copy_only=False | |
| ) | |
| local_indices = np.flatnonzero(null_mask)[:3] | |
| global_indices = ( | |
| local_indices + batch_start | |
| ).tolist() | |
| raise ValueError( | |
| f"Duration column '{DURATION_COLUMN}' contains null " | |
| f"values at rows {global_indices}." | |
| ) | |
| try: | |
| float_batch = pc.cast( | |
| arrow_batch, | |
| pa.float64(), | |
| safe=False, | |
| ) | |
| except ( | |
| pa.ArrowInvalid, | |
| pa.ArrowNotImplementedError, | |
| TypeError, | |
| ) as exc: | |
| raise TypeError( | |
| f"Duration column '{DURATION_COLUMN}' cannot be cast " | |
| f"to float64 from Arrow type {arrow_batch.type}." | |
| ) from exc | |
| duration_seconds = float_batch.to_numpy( | |
| zero_copy_only=False | |
| ) | |
| invalid_mask = ( | |
| ~np.isfinite(duration_seconds) | |
| | (duration_seconds <= 0.0) | |
| ) | |
| if invalid_mask.any(): | |
| local_indices = np.flatnonzero(invalid_mask)[:3] | |
| global_indices = ( | |
| local_indices + batch_start | |
| ).tolist() | |
| invalid_values = duration_seconds[ | |
| local_indices | |
| ].tolist() | |
| raise ValueError( | |
| f"Duration column '{DURATION_COLUMN}' contains invalid " | |
| f"values at rows {global_indices}: {invalid_values}." | |
| ) | |
| batch_max_seconds = float(duration_seconds.max()) | |
| if batch_max_seconds > MAX_AUDIO_DURATION_SECONDS: | |
| local_index = int(np.argmax(duration_seconds)) | |
| raise ValueError( | |
| f"Metadata duration column contains " | |
| f"{batch_max_seconds:.3f}s at row " | |
| f"{batch_start + local_index}, which exceeds " | |
| f"MAX_AUDIO_DURATION_SECONDS=" | |
| f"{MAX_AUDIO_DURATION_SECONDS}." | |
| ) | |
| duration_samples = np.ceil( | |
| duration_seconds * sample_rate | |
| ).astype(np.int64, copy=False) | |
| batch_max_samples = int(duration_samples.max()) | |
| observed_max_samples = max( | |
| observed_max_samples, | |
| batch_max_samples, | |
| ) | |
| bin_ids = ( | |
| duration_samples - 1 | |
| ) // histogram_bin_samples | |
| counts += np.bincount( | |
| bin_ids, | |
| minlength=histogram_size, | |
| ).astype(np.int64, copy=False) | |
| sample_sums += np.bincount( | |
| bin_ids, | |
| weights=duration_samples, | |
| minlength=histogram_size, | |
| ).astype(np.float64, copy=False) | |
| np.maximum.at( | |
| bin_maxima, | |
| bin_ids, | |
| duration_samples, | |
| ) | |
| processed_rows += batch_length | |
| progress.update(batch_length) | |
| progress.set_postfix( | |
| max_s=f"{observed_max_samples / sample_rate:.3f}", | |
| refresh=False, | |
| ) | |
| progress.close() | |
| if processed_rows != total_rows: | |
| raise RuntimeError( | |
| f"Arrow duration scan processed {processed_rows} rows, " | |
| f"but the metadata dataset contains {total_rows} rows." | |
| ) | |
| occupied = counts > 0 | |
| return ( | |
| counts[occupied], | |
| sample_sums[occupied], | |
| bin_maxima[occupied], | |
| observed_max_samples, | |
| ) | |
| def optimize_bucket_boundaries_for_count( | |
| counts, | |
| sample_sums, | |
| bin_maxima, | |
| bucket_count, | |
| max_allowed_samples, | |
| ): | |
| occupied_bins = len(counts) | |
| if bucket_count > occupied_bins: | |
| raise ValueError( | |
| f"Cannot create {bucket_count} non-empty duration buckets from " | |
| f"only {occupied_bins} occupied duration bins." | |
| ) | |
| prefix_counts = np.concatenate( | |
| ([0], np.cumsum(counts, dtype=np.int64)) | |
| ) | |
| prefix_sums = np.concatenate( | |
| ([0.0], np.cumsum(sample_sums, dtype=np.float64)) | |
| ) | |
| costs = np.full( | |
| (bucket_count + 1, occupied_bins), | |
| np.inf, | |
| dtype=np.float64, | |
| ) | |
| backpointers = np.full( | |
| (bucket_count + 1, occupied_bins), | |
| -1, | |
| dtype=np.int32, | |
| ) | |
| for end_index in range(occupied_bins): | |
| ceiling = ( | |
| max_allowed_samples | |
| if bucket_count == 1 and end_index == occupied_bins - 1 | |
| else int(bin_maxima[end_index]) | |
| ) | |
| segment_count = prefix_counts[end_index + 1] | |
| segment_sum = prefix_sums[end_index + 1] | |
| costs[1, end_index] = ( | |
| ceiling * segment_count - segment_sum | |
| ) | |
| for used_buckets in range(2, bucket_count + 1): | |
| for end_index in range(used_buckets - 1, occupied_bins): | |
| starts = np.arange( | |
| used_buckets - 1, | |
| end_index + 1, | |
| dtype=np.int64, | |
| ) | |
| previous_costs = costs[ | |
| used_buckets - 1, | |
| starts - 1, | |
| ] | |
| ceiling = ( | |
| max_allowed_samples | |
| if ( | |
| used_buckets == bucket_count | |
| and end_index == occupied_bins - 1 | |
| ) | |
| else int(bin_maxima[end_index]) | |
| ) | |
| segment_counts = ( | |
| prefix_counts[end_index + 1] | |
| - prefix_counts[starts] | |
| ) | |
| segment_sums = ( | |
| prefix_sums[end_index + 1] | |
| - prefix_sums[starts] | |
| ) | |
| candidate_costs = previous_costs + ( | |
| ceiling * segment_counts - segment_sums | |
| ) | |
| best_offset = int(np.argmin(candidate_costs)) | |
| costs[used_buckets, end_index] = candidate_costs[ | |
| best_offset | |
| ] | |
| backpointers[used_buckets, end_index] = int( | |
| starts[best_offset] | |
| ) | |
| end_index = occupied_bins - 1 | |
| endpoints = [] | |
| for used_buckets in range(bucket_count, 0, -1): | |
| endpoints.append(end_index) | |
| if used_buckets == 1: | |
| break | |
| start_index = int( | |
| backpointers[used_buckets, end_index] | |
| ) | |
| end_index = start_index - 1 | |
| endpoints.reverse() | |
| boundaries = [ | |
| int(bin_maxima[index]) | |
| for index in endpoints | |
| ] | |
| boundaries[-1] = max_allowed_samples | |
| return tuple(boundaries), float( | |
| costs[bucket_count, occupied_bins - 1] | |
| ) | |
| def calculate_duration_bucket_boundaries( | |
| metadata_dataset, | |
| sample_rate, | |
| ): | |
| histogram_bin_samples = max( | |
| 1, | |
| int( | |
| round( | |
| DURATION_HISTOGRAM_RESOLUTION_SECONDS | |
| * sample_rate | |
| ) | |
| ), | |
| ) | |
| ( | |
| counts, | |
| sample_sums, | |
| bin_maxima, | |
| max_allowed_samples, | |
| ) = build_metadata_duration_histogram( | |
| metadata_dataset=metadata_dataset, | |
| sample_rate=sample_rate, | |
| histogram_bin_samples=histogram_bin_samples, | |
| ) | |
| occupied_bins = len(counts) | |
| maximum_candidate_buckets = min( | |
| MAX_DURATION_BUCKETS, | |
| occupied_bins, | |
| ) | |
| minimum_candidate_buckets = min( | |
| MIN_DURATION_BUCKETS, | |
| maximum_candidate_buckets, | |
| ) | |
| candidates = [] | |
| actual_sample_total = float(sample_sums.sum()) | |
| bucket_counts_to_test = range( | |
| minimum_candidate_buckets, | |
| maximum_candidate_buckets + 1, | |
| ) | |
| optimization_progress = tqdm( | |
| bucket_counts_to_test, | |
| total=( | |
| maximum_candidate_buckets | |
| - minimum_candidate_buckets | |
| + 1 | |
| ), | |
| desc="Optimizing duration buckets", | |
| unit="candidate", | |
| dynamic_ncols=True, | |
| ) | |
| for bucket_count in optimization_progress: | |
| boundaries, padding_cost = ( | |
| optimize_bucket_boundaries_for_count( | |
| counts=counts, | |
| sample_sums=sample_sums, | |
| bin_maxima=bin_maxima, | |
| bucket_count=bucket_count, | |
| max_allowed_samples=max_allowed_samples, | |
| ) | |
| ) | |
| padding_ratio = padding_cost / actual_sample_total | |
| candidates.append( | |
| ( | |
| bucket_count, | |
| boundaries, | |
| padding_cost, | |
| padding_ratio, | |
| ) | |
| ) | |
| best_padding_cost = min( | |
| candidate[2] | |
| for candidate in candidates | |
| ) | |
| allowed_padding_cost = best_padding_cost * ( | |
| 1.0 + DURATION_BUCKET_ELBOW_TOLERANCE | |
| ) | |
| selected = next( | |
| candidate | |
| for candidate in candidates | |
| if candidate[2] <= allowed_padding_cost | |
| ) | |
| print("Duration bucket optimization:") | |
| for bucket_count, _, _, padding_ratio in candidates: | |
| marker = " *" if bucket_count == selected[0] else "" | |
| print( | |
| f" {bucket_count:2d} buckets: " | |
| f"estimated padding {padding_ratio * 100:.3f}%{marker}" | |
| ) | |
| boundaries_samples = selected[1] | |
| boundaries_seconds = tuple( | |
| boundary / sample_rate | |
| for boundary in boundaries_samples | |
| ) | |
| histogram_bucket_assignments = np.searchsorted( | |
| np.asarray(boundaries_samples, dtype=np.int64), | |
| bin_maxima, | |
| side="left", | |
| ) | |
| bucket_counts = np.bincount( | |
| histogram_bucket_assignments, | |
| weights=counts, | |
| minlength=len(boundaries_samples), | |
| ).astype(np.int64, copy=False) | |
| print( | |
| f"Selected {len(boundaries_samples)} duration buckets from " | |
| f"all {int(counts.sum()):,} metadata durations." | |
| ) | |
| for index, (ceiling, count) in enumerate( | |
| zip(boundaries_seconds, bucket_counts), | |
| start=1, | |
| ): | |
| print( | |
| f" Bucket {index:02d}: <= {ceiling:.3f}s " | |
| f"({int(count):,} metadata rows)" | |
| ) | |
| return boundaries_samples, boundaries_seconds | |
| def get_metadata_rows_by_indices(metadata_dataset, row_indices): | |
| """Read a small random-access metadata batch and return row dictionaries.""" | |
| indices = [int(index) for index in row_indices] | |
| columnar_batch = metadata_dataset[indices] | |
| return [ | |
| { | |
| column_name: columnar_batch[column_name][position] | |
| for column_name in metadata_dataset.column_names | |
| } | |
| for position in range(len(indices)) | |
| ] | |
| def build_duration_bucket_indices( | |
| metadata_dataset, | |
| sample_rate, | |
| bucket_boundaries_samples, | |
| ): | |
| duration_column = metadata_dataset.select_columns( | |
| [DURATION_COLUMN] | |
| ).data.column(DURATION_COLUMN) | |
| total_rows = len(metadata_dataset) | |
| boundaries = np.asarray( | |
| bucket_boundaries_samples, | |
| dtype=np.int64, | |
| ) | |
| index_chunks_by_bucket = [ | |
| [] | |
| for _ in bucket_boundaries_samples | |
| ] | |
| bucket_counts = np.zeros( | |
| len(bucket_boundaries_samples), | |
| dtype=np.int64, | |
| ) | |
| progress = tqdm( | |
| total=total_rows, | |
| desc="Assigning global duration buckets", | |
| unit="rows", | |
| dynamic_ncols=True, | |
| ) | |
| for batch_start in range( | |
| 0, | |
| total_rows, | |
| DURATION_BUCKET_INDEX_SCAN_BATCH_SIZE, | |
| ): | |
| batch_length = min( | |
| DURATION_BUCKET_INDEX_SCAN_BATCH_SIZE, | |
| total_rows - batch_start, | |
| ) | |
| arrow_batch = duration_column.slice( | |
| batch_start, | |
| batch_length, | |
| ) | |
| if isinstance(arrow_batch, pa.ChunkedArray): | |
| arrow_batch = arrow_batch.combine_chunks() | |
| if arrow_batch.null_count: | |
| raise ValueError( | |
| f"Duration column '{DURATION_COLUMN}' contains null values " | |
| f"while assigning global buckets near row {batch_start}." | |
| ) | |
| duration_seconds = pc.cast( | |
| arrow_batch, | |
| pa.float64(), | |
| safe=False, | |
| ).to_numpy(zero_copy_only=False) | |
| duration_samples = np.ceil( | |
| duration_seconds * sample_rate | |
| ).astype(np.int64, copy=False) | |
| assignments = np.searchsorted( | |
| boundaries, | |
| duration_samples, | |
| side="left", | |
| ) | |
| invalid_mask = assignments >= len(boundaries) | |
| if invalid_mask.any(): | |
| local_index = int(np.flatnonzero(invalid_mask)[0]) | |
| raise ValueError( | |
| f"Duration at metadata row {batch_start + local_index} " | |
| f"exceeds the final optimized bucket ceiling." | |
| ) | |
| global_indices = np.arange( | |
| batch_start, | |
| batch_start + batch_length, | |
| dtype=np.int64, | |
| ) | |
| for bucket_index in np.unique(assignments): | |
| bucket_mask = assignments == bucket_index | |
| bucket_indices = global_indices[bucket_mask] | |
| index_chunks_by_bucket[int(bucket_index)].append( | |
| bucket_indices | |
| ) | |
| bucket_counts[int(bucket_index)] += len(bucket_indices) | |
| progress.update(batch_length) | |
| progress.set_postfix( | |
| assigned=batch_start + batch_length, | |
| refresh=False, | |
| ) | |
| progress.close() | |
| bucket_indices = [] | |
| for chunks in index_chunks_by_bucket: | |
| if chunks: | |
| bucket_indices.append(np.concatenate(chunks)) | |
| else: | |
| bucket_indices.append(np.empty(0, dtype=np.int64)) | |
| if int(bucket_counts.sum()) != total_rows: | |
| raise RuntimeError( | |
| f"Assigned {int(bucket_counts.sum())} rows to duration buckets, " | |
| f"but the metadata dataset contains {total_rows} rows." | |
| ) | |
| print("Global duration bucket membership:") | |
| for bucket_index, (boundary, count) in enumerate( | |
| zip(bucket_boundaries_samples, bucket_counts), | |
| start=1, | |
| ): | |
| print( | |
| f" Bucket {bucket_index:02d}: <= " | |
| f"{boundary / sample_rate:.3f}s " | |
| f"({int(count):,} rows, " | |
| f"{math.ceil(int(count) / BATCH_SIZE):,} batches)" | |
| ) | |
| return bucket_indices, bucket_counts | |
| def prepare_global_bucket_batch( | |
| rows, | |
| audio_pool, | |
| audio_datasets, | |
| audio_key_index, | |
| sample_rate, | |
| padded_length, | |
| bucket_upper_seconds, | |
| ): | |
| items = ( | |
| ( | |
| position, | |
| row, | |
| audio_datasets, | |
| audio_key_index, | |
| sample_rate, | |
| ) | |
| for position, row in enumerate(rows) | |
| ) | |
| prepared = list(audio_pool.map(prepare_audio_item, items)) | |
| batch_rows = [None] * len(prepared) | |
| errors = [None] * len(prepared) | |
| audio_by_position = {} | |
| for position, row, audio, error in prepared: | |
| batch_rows[position] = row | |
| errors[position] = error | |
| if audio is None: | |
| continue | |
| if len(audio) > padded_length: | |
| errors[position] = ( | |
| f"ERROR: decoded audio length is " | |
| f"{len(audio) / sample_rate:.3f}s, which exceeds its " | |
| f"metadata-assigned bucket ceiling of " | |
| f"{bucket_upper_seconds:.3f}s" | |
| ) | |
| continue | |
| audio_by_position[position] = audio | |
| valid_positions = sorted(audio_by_position) | |
| if not valid_positions: | |
| unique_errors = [] | |
| for error in errors: | |
| if error is not None and error not in unique_errors: | |
| unique_errors.append(error) | |
| if len(unique_errors) == 3: | |
| break | |
| details = " | ".join(unique_errors) or "No error detail was returned." | |
| raise RuntimeError( | |
| "Every row in the global duration-bucket batch failed before " | |
| f"codec extraction. First errors: {details}" | |
| ) | |
| audio_tensor, audio_lens = build_padded_audio_batch( | |
| audio_by_position=audio_by_position, | |
| positions=valid_positions, | |
| padded_length=padded_length, | |
| ) | |
| return { | |
| "rows": batch_rows, | |
| "errors": errors, | |
| "valid_positions": valid_positions, | |
| "audio_tensor": audio_tensor, | |
| "audio_lens": audio_lens, | |
| } | |
| def build_padded_audio_batch( | |
| audio_by_position, | |
| positions, | |
| padded_length, | |
| ): | |
| audio_list = [audio_by_position[position] for position in positions] | |
| original_lengths = [len(audio) for audio in audio_list] | |
| padded_batch = [] | |
| for audio in audio_list: | |
| pad_len = padded_length - len(audio) | |
| if pad_len < 0: | |
| raise ValueError( | |
| f"Audio length {len(audio)} exceeds its bucket padding " | |
| f"length {padded_length}." | |
| ) | |
| if pad_len > 0: | |
| audio = F.pad( | |
| audio, | |
| (0, pad_len), | |
| mode="constant", | |
| value=0, | |
| ) | |
| padded_batch.append(audio) | |
| return ( | |
| torch.stack(padded_batch, dim=0), | |
| torch.tensor(original_lengths, dtype=torch.long), | |
| ) | |
| def encode_batch(model, pre_quantized, latent_len): | |
| encoded_tokens = model.codec.vector_quantizer.encode( | |
| inputs=pre_quantized, | |
| input_len=latent_len, | |
| ) | |
| trimmed_codes = [] | |
| for idx in range(pre_quantized.shape[0]): | |
| tlen = int(latent_len[idx].item()) | |
| trimmed_codes.append(encoded_tokens[:, idx, :tlen]) | |
| return trimmed_codes | |
| def get_audio_encoder(model): | |
| if hasattr(model, "audio_encoder"): | |
| return model.audio_encoder | |
| if hasattr(model, "codec") and hasattr(model.codec, "audio_encoder"): | |
| return model.codec.audio_encoder | |
| raise AttributeError("Could not find audio_encoder on codec model.") | |
| def parse_audio_encoder_output(enc_out, batch_size): | |
| pre_quantized = None | |
| latent_len = None | |
| if torch.is_tensor(enc_out): | |
| pre_quantized = enc_out | |
| elif isinstance(enc_out, tuple): | |
| pre_quantized = enc_out[0] | |
| for item in enc_out[1:]: | |
| if ( | |
| torch.is_tensor(item) | |
| and item.dim() == 1 | |
| and item.numel() == batch_size | |
| ): | |
| latent_len = item | |
| break | |
| elif isinstance(enc_out, dict): | |
| pre_quantized = ( | |
| enc_out.get("encoded") | |
| or enc_out.get("latents") | |
| or enc_out.get("encoder_out") | |
| or enc_out.get("pre_quantized") | |
| ) | |
| latent_len = ( | |
| enc_out.get("encoded_len") | |
| or enc_out.get("latents_len") | |
| or enc_out.get("encoder_out_len") | |
| or enc_out.get("lengths") | |
| or enc_out.get("audio_codes_len") | |
| ) | |
| else: | |
| raise ValueError( | |
| f"Unsupported audio_encoder output type: {type(enc_out)}" | |
| ) | |
| if pre_quantized is None: | |
| raise ValueError( | |
| "Could not locate pre_quantized latents in audio_encoder output." | |
| ) | |
| return pre_quantized, latent_len | |
| def extract_prequant_batch(model, wav_batch, audio_lens, device): | |
| wav_batch = wav_batch.to( | |
| device=device, | |
| dtype=INFERENCE_DTYPE, | |
| non_blocking=True, | |
| ) | |
| audio_lens = audio_lens.to(device, non_blocking=True) | |
| encoder = get_audio_encoder(model) | |
| enc_out = encoder(audio=wav_batch, audio_len=audio_lens) | |
| pre_quantized, latent_len = parse_audio_encoder_output( | |
| enc_out, | |
| wav_batch.shape[0], | |
| ) | |
| if pre_quantized.dim() != 3: | |
| raise ValueError( | |
| f"Expected 3D pre_quantized tensor, got shape " | |
| f"{tuple(pre_quantized.shape)}" | |
| ) | |
| if latent_len is None: | |
| raise ValueError( | |
| "The codec audio encoder did not return latent lengths." | |
| ) | |
| time_major_latents = pre_quantized.transpose(1, 2) | |
| trimmed_latents = [] | |
| for idx in range(time_major_latents.shape[0]): | |
| tlen = int(latent_len[idx].item()) | |
| trimmed_latents.append(time_major_latents[idx, :tlen, :]) | |
| return trimmed_latents, pre_quantized, latent_len | |
| #checkpoint | |
| def build_checkpoint_features(metadata_dataset): | |
| generated_columns = {ERROR_COLUMN} | |
| if SAVE_PREQUANT: | |
| generated_columns.add(PREQUANT_COLUMN) | |
| if SAVE_DISCRETE: | |
| generated_columns.add(DISCRETE_TOKENS_COLUMN) | |
| conflicting_columns = generated_columns.intersection( | |
| metadata_dataset.column_names | |
| ) | |
| if conflicting_columns: | |
| raise KeyError( | |
| f"Primary dataset already contains generated output columns: " | |
| f"{sorted(conflicting_columns)}" | |
| ) | |
| features = dict(metadata_dataset.features) | |
| features.pop(AUDIO_COLUMN, None) | |
| if SAVE_PREQUANT: | |
| features[PREQUANT_COLUMN] = Sequence( | |
| Sequence(Value("float16")) | |
| ) | |
| if SAVE_DISCRETE: | |
| features[DISCRETE_TOKENS_COLUMN] = Sequence( | |
| Sequence(Value("int64")) | |
| ) | |
| features[ERROR_COLUMN] = Value("string") | |
| return Features(features) | |
| def find_saved_state(): | |
| checkpoint_root = Path(CHECKPOINT_DIR) | |
| checkpoint_root.mkdir(parents=True, exist_ok=True) | |
| shard_paths = sorted( | |
| path | |
| for path in checkpoint_root.glob("shard-*") | |
| if path.is_dir() | |
| ) | |
| processed_rows = 0 | |
| checkpoint_progress = tqdm( | |
| shard_paths, | |
| total=len(shard_paths), | |
| desc="Scanning checkpoint shards", | |
| unit="shard", | |
| dynamic_ncols=True, | |
| disable=not shard_paths, | |
| ) | |
| for path in checkpoint_progress: | |
| processed_rows += len(load_from_disk(str(path))) | |
| checkpoint_progress.set_postfix( | |
| rows=processed_rows, | |
| refresh=False, | |
| ) | |
| state_path = checkpoint_root / "state.json" | |
| if not state_path.exists(): | |
| if shard_paths: | |
| raise RuntimeError( | |
| "Checkpoint shards exist without state.json. Global bucket " | |
| "resume requires both; clear CHECKPOINT_DIR for a fresh run." | |
| ) | |
| return 0, 0, 0, None, 0, 0 | |
| with state_path.open("r", encoding="utf-8") as handle: | |
| state = json.load(handle) | |
| if state.get("processing_order") != "global_duration_buckets": | |
| raise RuntimeError( | |
| "The checkpoint state was not created by the global " | |
| "duration-bucket pipeline. Clear CHECKPOINT_DIR before running." | |
| ) | |
| saved_processed_rows = int(state["processed_rows"]) | |
| if saved_processed_rows != processed_rows: | |
| raise RuntimeError( | |
| f"state.json reports {saved_processed_rows} processed rows, but " | |
| f"checkpoint shards contain {processed_rows}." | |
| ) | |
| saved_bucket_boundaries = tuple( | |
| int(boundary) | |
| for boundary in state["duration_bucket_boundaries_samples"] | |
| ) | |
| return ( | |
| len(shard_paths), | |
| processed_rows, | |
| int(state["processed_steps"]), | |
| saved_bucket_boundaries, | |
| int(state["current_bucket_index"]), | |
| int(state["bucket_offset"]), | |
| ) | |
| def save_state( | |
| processed_rows, | |
| processed_steps, | |
| duration_bucket_boundaries_samples, | |
| current_bucket_index, | |
| bucket_offset, | |
| ): | |
| checkpoint_root = Path(CHECKPOINT_DIR) | |
| state_path = checkpoint_root / "state.json" | |
| temp_path = checkpoint_root / ".state.json.tmp" | |
| with temp_path.open("w", encoding="utf-8") as handle: | |
| json.dump( | |
| { | |
| "processing_order": "global_duration_buckets", | |
| "processed_rows": processed_rows, | |
| "processed_steps": processed_steps, | |
| "current_bucket_index": current_bucket_index, | |
| "bucket_offset": bucket_offset, | |
| "duration_bucket_boundaries_samples": list( | |
| duration_bucket_boundaries_samples | |
| ), | |
| "duration_bucket_source": f"metadata:{DURATION_COLUMN}", | |
| }, | |
| handle, | |
| ) | |
| os.replace(temp_path, state_path) | |
| def save_shard( | |
| rows, | |
| shard_index, | |
| processed_rows, | |
| processed_steps, | |
| checkpoint_features, | |
| duration_bucket_boundaries_samples, | |
| current_bucket_index, | |
| bucket_offset, | |
| ): | |
| if not rows: | |
| return shard_index | |
| checkpoint_root = Path(CHECKPOINT_DIR) | |
| final_path = checkpoint_root / f"shard-{shard_index:06d}" | |
| temp_path = checkpoint_root / f".shard-{shard_index:06d}.tmp" | |
| if final_path.exists(): | |
| raise FileExistsError( | |
| f"Checkpoint shard already exists: {final_path}" | |
| ) | |
| if temp_path.exists(): | |
| shutil.rmtree(temp_path) | |
| Dataset.from_list( | |
| rows, | |
| features=checkpoint_features, | |
| ).save_to_disk( | |
| str(temp_path), | |
| max_shard_size=HF_MAX_SHARD_SIZE, | |
| ) | |
| os.replace(temp_path, final_path) | |
| save_state( | |
| processed_rows=processed_rows, | |
| processed_steps=processed_steps, | |
| duration_bucket_boundaries_samples=( | |
| duration_bucket_boundaries_samples | |
| ), | |
| current_bucket_index=current_bucket_index, | |
| bucket_offset=bucket_offset, | |
| ) | |
| print(f"Saved {len(rows)} rows to {final_path}") | |
| return shard_index + 1 | |
| def process_global_bucket_batch( | |
| model, | |
| batch, | |
| device, | |
| bucket_upper_seconds, | |
| ): | |
| rows = batch["rows"] | |
| errors = list(batch["errors"]) | |
| valid_positions = batch["valid_positions"] | |
| discrete_values = [None] * len(rows) | |
| prequant_values = [None] * len(rows) | |
| dim_skipped = 0 | |
| try: | |
| with torch.autocast( | |
| device_type="cuda", | |
| dtype=INFERENCE_DTYPE, | |
| ): | |
| ( | |
| trimmed_latents, | |
| pre_quantized, | |
| latent_len, | |
| ) = extract_prequant_batch( | |
| model, | |
| batch["audio_tensor"], | |
| batch["audio_lens"], | |
| device, | |
| ) | |
| trimmed_codes = None | |
| if SAVE_DISCRETE: | |
| trimmed_codes = encode_batch( | |
| model, | |
| pre_quantized, | |
| latent_len, | |
| ) | |
| for model_position, row_position in enumerate(valid_positions): | |
| if SAVE_PREQUANT and trimmed_latents is not None: | |
| latent = trimmed_latents[model_position] | |
| if latent.shape[-1] != EXPECTED_PREQUANT_DIM: | |
| errors[row_position] = ( | |
| f"ERROR: prequant dim is {latent.shape[-1]}, " | |
| f"expected {EXPECTED_PREQUANT_DIM}" | |
| ) | |
| dim_skipped += 1 | |
| continue | |
| prequant_values[row_position] = ( | |
| latent.cpu() | |
| .to(torch.float16) | |
| .numpy() | |
| .tolist() | |
| ) | |
| if SAVE_DISCRETE and trimmed_codes is not None: | |
| discrete_values[row_position] = ( | |
| trimmed_codes[model_position] | |
| .cpu() | |
| .to(torch.int64) | |
| .numpy() | |
| .tolist() | |
| ) | |
| except Exception as exc: | |
| raise RuntimeError( | |
| f"Codec extraction failed for the " | |
| f"{bucket_upper_seconds:.3f}s global duration bucket: {exc}" | |
| ) from exc | |
| processed_rows = [] | |
| successful_rows = 0 | |
| error_rows = 0 | |
| for position, row in enumerate(rows): | |
| output_row = dict(row) | |
| output_row.pop(AUDIO_COLUMN, None) | |
| if SAVE_PREQUANT: | |
| output_row[PREQUANT_COLUMN] = prequant_values[position] | |
| if SAVE_DISCRETE: | |
| output_row[DISCRETE_TOKENS_COLUMN] = discrete_values[position] | |
| output_row[ERROR_COLUMN] = errors[position] | |
| if errors[position] is None: | |
| successful_rows += 1 | |
| else: | |
| error_rows += 1 | |
| processed_rows.append(output_row) | |
| return processed_rows, successful_rows, error_rows, dim_skipped | |
| def main(): | |
| print("=" * 60) | |
| print("Unified Codec Extraction Script") | |
| print(f" Metadata dataset: {DATASET_SOURCE}") | |
| print(f" Audio dataset: {AUDIO_DATASET_SOURCE}") | |
| print(f" Checkpoint dir: {CHECKPOINT_DIR}") | |
| print(f" Save discrete: {SAVE_DISCRETE}") | |
| print(f" Save prequant: {SAVE_PREQUANT}") | |
| print(f" Batch size: {BATCH_SIZE}") | |
| print( | |
| f" Duration buckets: automatic " | |
| f"({MIN_DURATION_BUCKETS}-{MAX_DURATION_BUCKETS})" | |
| ) | |
| print(f" Processing order: global bucket-by-bucket") | |
| print(f" Duration column: {DURATION_COLUMN}") | |
| print(f" Maximum duration: {MAX_AUDIO_DURATION_SECONDS:.1f}s") | |
| print(f" Workers: {NUM_WORKERS}") | |
| print(f" Inference dtype: {INFERENCE_DTYPE}") | |
| print("=" * 60) | |
| if not SAVE_DISCRETE and not SAVE_PREQUANT: | |
| print( | |
| "ERROR: Both SAVE_DISCRETE and SAVE_PREQUANT are False. " | |
| "Nothing to do." | |
| ) | |
| sys.exit(1) | |
| print("Loading metadata dataset...") | |
| metadata_dataset = load_input_dataset() | |
| print(f"Loaded {len(metadata_dataset):,} metadata rows.") | |
| print("Building checkpoint schema...") | |
| checkpoint_features = build_checkpoint_features(metadata_dataset) | |
| print("Checkpoint schema ready.") | |
| total_rows = len(metadata_dataset) | |
| ( | |
| shard_index, | |
| processed_count, | |
| processed_steps, | |
| saved_bucket_boundaries_samples, | |
| current_bucket_index, | |
| bucket_offset, | |
| ) = find_saved_state() | |
| print( | |
| f"Checkpoint state: {processed_count:,}/{total_rows:,} " | |
| f"rows completed across {shard_index:,} shards." | |
| ) | |
| if saved_bucket_boundaries_samples is None: | |
| print( | |
| "Calculating duration buckets from the Arrow-backed " | |
| f"'{DURATION_COLUMN}' column..." | |
| ) | |
| ( | |
| duration_bucket_boundaries_samples, | |
| duration_bucket_boundaries_seconds, | |
| ) = calculate_duration_bucket_boundaries( | |
| metadata_dataset=metadata_dataset, | |
| sample_rate=TARGET_SAMPLE_RATE, | |
| ) | |
| else: | |
| duration_bucket_boundaries_samples = ( | |
| saved_bucket_boundaries_samples | |
| ) | |
| duration_bucket_boundaries_seconds = tuple( | |
| boundary / TARGET_SAMPLE_RATE | |
| for boundary in duration_bucket_boundaries_samples | |
| ) | |
| print( | |
| "Reusing duration buckets saved in state.json: " | |
| f"{tuple(round(value, 3) for value in duration_bucket_boundaries_seconds)}" | |
| ) | |
| print("Building global metadata row indices for each duration bucket...") | |
| duration_bucket_indices, duration_bucket_counts = ( | |
| build_duration_bucket_indices( | |
| metadata_dataset=metadata_dataset, | |
| sample_rate=TARGET_SAMPLE_RATE, | |
| bucket_boundaries_samples=( | |
| duration_bucket_boundaries_samples | |
| ), | |
| ) | |
| ) | |
| if current_bucket_index > len(duration_bucket_indices): | |
| raise ValueError( | |
| f"Checkpoint current_bucket_index={current_bucket_index} " | |
| f"exceeds the {len(duration_bucket_indices)} buckets." | |
| ) | |
| if current_bucket_index < len(duration_bucket_indices): | |
| if bucket_offset > len( | |
| duration_bucket_indices[current_bucket_index] | |
| ): | |
| raise ValueError( | |
| f"Checkpoint bucket_offset={bucket_offset} exceeds " | |
| f"bucket {current_bucket_index + 1} size." | |
| ) | |
| elif bucket_offset != 0: | |
| raise ValueError( | |
| "A completed checkpoint must have bucket_offset=0." | |
| ) | |
| expected_processed_rows = sum( | |
| len(duration_bucket_indices[index]) | |
| for index in range(current_bucket_index) | |
| ) | |
| if current_bucket_index < len(duration_bucket_indices): | |
| expected_processed_rows += bucket_offset | |
| if expected_processed_rows != processed_count: | |
| raise RuntimeError( | |
| f"Checkpoint cursor implies {expected_processed_rows} processed " | |
| f"rows, but checkpoint shards contain {processed_count}." | |
| ) | |
| print(f"Loading audio datasets from: {AUDIO_DATASET_SOURCE}") | |
| audio_datasets = load_audio_datasets() | |
| total_audio_rows = sum(len(dataset) for dataset in audio_datasets) | |
| print( | |
| f"Building audio lookup from {total_audio_rows:,} audio rows " | |
| f"across {len(audio_datasets)} datasets." | |
| ) | |
| audio_key_index = build_audio_key_index(audio_datasets) | |
| print(f"Indexed {len(audio_key_index):,} unique audio keys.") | |
| if not torch.cuda.is_available(): | |
| raise RuntimeError("CUDA is required for bfloat16 codec inference.") | |
| if not torch.cuda.is_bf16_supported(): | |
| raise RuntimeError( | |
| "The active CUDA device does not support bfloat16 inference." | |
| ) | |
| device = "cuda" | |
| print(f"Loading Kanadec nano audio tokenizer on {device}...") | |
| codec_model = load_dune_audio_tokenizer( | |
| MODEL_ID, | |
| device=device, | |
| ) | |
| codec_model = codec_model.to( | |
| device=device, | |
| dtype=INFERENCE_DTYPE, | |
| ).eval() | |
| codec_model.codec.audio_encoder = torch.compile(codec_model.codec.audio_encoder, mode="max-autotune-no-cudagraphs") | |
| floating_dtypes = { | |
| tensor.dtype | |
| for tensor in ( | |
| *codec_model.parameters(), | |
| *codec_model.buffers(), | |
| ) | |
| if tensor.is_floating_point() | |
| } | |
| if floating_dtypes != {INFERENCE_DTYPE}: | |
| raise RuntimeError( | |
| f"Codec floating dtypes are {sorted(map(str, floating_dtypes))}, " | |
| f"expected only {INFERENCE_DTYPE}." | |
| ) | |
| print(f"Target Sample Rate: {TARGET_SAMPLE_RATE}") | |
| print(f"Confirmed model dtype: {INFERENCE_DTYPE}") | |
| total_progress_steps = sum( | |
| math.ceil(int(count) / BATCH_SIZE) | |
| for count in duration_bucket_counts | |
| ) | |
| progress = tqdm( | |
| total=total_progress_steps, | |
| initial=processed_steps, | |
| desc="Extracting codec features", | |
| unit="step", | |
| dynamic_ncols=True, | |
| ) | |
| pending_rows = [] | |
| steps_since_checkpoint = 0 | |
| total_successful = 0 | |
| total_errors = 0 | |
| total_dim_skipped = 0 | |
| started_at = time.time() | |
| try: | |
| with ThreadPoolExecutor(max_workers=NUM_WORKERS) as audio_pool: | |
| with torch.no_grad(): | |
| for bucket_index in range( | |
| current_bucket_index, | |
| len(duration_bucket_indices), | |
| ): | |
| row_indices = duration_bucket_indices[bucket_index] | |
| start_offset = ( | |
| bucket_offset | |
| if bucket_index == current_bucket_index | |
| else 0 | |
| ) | |
| bucket_upper_samples = ( | |
| duration_bucket_boundaries_samples[bucket_index] | |
| ) | |
| bucket_upper_seconds = ( | |
| duration_bucket_boundaries_seconds[bucket_index] | |
| ) | |
| print( | |
| f"Processing bucket {bucket_index + 1}/" | |
| f"{len(duration_bucket_indices)}: <= " | |
| f"{bucket_upper_seconds:.3f}s, " | |
| f"{len(row_indices):,} rows, starting at " | |
| f"offset {start_offset:,}." | |
| ) | |
| for batch_start in range( | |
| start_offset, | |
| len(row_indices), | |
| BATCH_SIZE, | |
| ): | |
| batch_indices = row_indices[ | |
| batch_start:batch_start + BATCH_SIZE | |
| ] | |
| rows = get_metadata_rows_by_indices( | |
| metadata_dataset, | |
| batch_indices, | |
| ) | |
| batch = prepare_global_bucket_batch( | |
| rows=rows, | |
| audio_pool=audio_pool, | |
| audio_datasets=audio_datasets, | |
| audio_key_index=audio_key_index, | |
| sample_rate=TARGET_SAMPLE_RATE, | |
| padded_length=bucket_upper_samples, | |
| bucket_upper_seconds=bucket_upper_seconds, | |
| ) | |
| ( | |
| processed_batch, | |
| successful_rows, | |
| error_rows, | |
| dim_skipped, | |
| ) = process_global_bucket_batch( | |
| model=codec_model, | |
| batch=batch, | |
| device=device, | |
| bucket_upper_seconds=bucket_upper_seconds, | |
| ) | |
| pending_rows.extend(processed_batch) | |
| processed_count += len(processed_batch) | |
| processed_steps += 1 | |
| steps_since_checkpoint += 1 | |
| total_successful += successful_rows | |
| total_errors += error_rows | |
| total_dim_skipped += dim_skipped | |
| next_offset = batch_start + len(batch_indices) | |
| if next_offset >= len(row_indices): | |
| current_bucket_index = bucket_index + 1 | |
| bucket_offset = 0 | |
| else: | |
| current_bucket_index = bucket_index | |
| bucket_offset = next_offset | |
| progress.update(1) | |
| progress.set_postfix( | |
| bucket=( | |
| f"{bucket_index + 1}/" | |
| f"{len(duration_bucket_indices)}" | |
| ), | |
| bucket_rows=( | |
| f"{min(next_offset, len(row_indices))}/" | |
| f"{len(row_indices)}" | |
| ), | |
| rows=processed_count, | |
| errors=total_errors, | |
| dim_skipped=total_dim_skipped, | |
| checkpoint_in=max( | |
| 0, | |
| CHECKPOINT_INTERVAL_STEPS | |
| - steps_since_checkpoint, | |
| ), | |
| refresh=False, | |
| ) | |
| if ( | |
| steps_since_checkpoint | |
| >= CHECKPOINT_INTERVAL_STEPS | |
| ): | |
| shard_index = save_shard( | |
| rows=pending_rows, | |
| shard_index=shard_index, | |
| processed_rows=processed_count, | |
| processed_steps=processed_steps, | |
| checkpoint_features=checkpoint_features, | |
| duration_bucket_boundaries_samples=( | |
| duration_bucket_boundaries_samples | |
| ), | |
| current_bucket_index=( | |
| current_bucket_index | |
| ), | |
| bucket_offset=bucket_offset, | |
| ) | |
| pending_rows = [] | |
| steps_since_checkpoint = 0 | |
| except KeyboardInterrupt: | |
| print("\nInterrupted; saving completed rows before exit.") | |
| finally: | |
| shard_index = save_shard( | |
| rows=pending_rows, | |
| shard_index=shard_index, | |
| processed_rows=processed_count, | |
| processed_steps=processed_steps, | |
| checkpoint_features=checkpoint_features, | |
| duration_bucket_boundaries_samples=( | |
| duration_bucket_boundaries_samples | |
| ), | |
| current_bucket_index=current_bucket_index, | |
| bucket_offset=bucket_offset, | |
| ) | |
| progress.close() | |
| elapsed = time.time() - started_at | |
| print("=" * 60) | |
| print("Done.") | |
| print(f" Successful: {total_successful}") | |
| print(f" Errors: {total_errors}") | |
| print(f" Dim-skipped: {total_dim_skipped}") | |
| print(f" Elapsed: {elapsed:.2f} seconds") | |
| print(f" Shards: {CHECKPOINT_DIR}") | |
| print("=" * 60) | |
| if __name__ == "__main__": | |
| main() |