| 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" |
|
|
|
|
| |
| DATASET_SOURCE = "/home/ubuntu/data" |
| DATASET_CONFIG = None |
| DATASET_SPLIT = "train" |
|
|
|
|
| |
| AUDIO_DATASET_SOURCE = ["/home/ubuntu/data"] |
| AUDIO_DATASET_CONFIG = None |
| AUDIO_DATASET_SPLIT = "train" |
|
|
| KEY_COLUMN = "key" |
| AUDIO_COLUMN = "audio" |
|
|
| |
| 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" |
|
|
| |
| |
| SAVE_PREQUANT = True |
| SAVE_DISCRETE = True |
|
|
| TEST_RUN = 0 |
|
|
| BATCH_SIZE = 64 |
| NUM_WORKERS = 32 |
| TARGET_SAMPLE_RATE = 22050 |
| INFERENCE_DTYPE = torch.bfloat16 |
|
|
|
|
| MAX_AUDIO_DURATION_SECONDS = 30.0 |
|
|
|
|
|
|
| |
| 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 |
|
|
| |
| 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), |
| ) |
|
|
|
|
| @torch.no_grad() |
| 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 |
|
|
|
|
| @torch.no_grad() |
| 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 |
|
|
|
|
| |
| 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() |