Update dune_extraction.py
Browse files- dune_extraction.py +13 -31
dune_extraction.py
CHANGED
|
@@ -20,16 +20,7 @@ import pyarrow as pa
|
|
| 20 |
import pyarrow.compute as pc
|
| 21 |
import torch
|
| 22 |
import torch.nn.functional as F
|
| 23 |
-
from datasets import
|
| 24 |
-
Audio,
|
| 25 |
-
Dataset,
|
| 26 |
-
DatasetDict,
|
| 27 |
-
Features,
|
| 28 |
-
Sequence,
|
| 29 |
-
Value,
|
| 30 |
-
load_dataset,
|
| 31 |
-
load_from_disk,
|
| 32 |
-
)
|
| 33 |
from tqdm.auto import tqdm
|
| 34 |
|
| 35 |
warnings.filterwarnings("ignore")
|
|
@@ -38,20 +29,21 @@ warnings.filterwarnings("ignore")
|
|
| 38 |
MODEL_ID = "Respair/dune_codec"
|
| 39 |
|
| 40 |
|
| 41 |
-
#
|
| 42 |
-
# Contains the metadata rows to enrich and a unique bridge key.
|
| 43 |
DATASET_SOURCE = "/home/ubuntu/data"
|
| 44 |
DATASET_CONFIG = None
|
| 45 |
DATASET_SPLIT = "train"
|
| 46 |
|
| 47 |
-
|
| 48 |
-
#
|
| 49 |
AUDIO_DATASET_SOURCE = ["/home/ubuntu/data"]
|
| 50 |
AUDIO_DATASET_CONFIG = None
|
| 51 |
AUDIO_DATASET_SPLIT = "train"
|
| 52 |
|
| 53 |
KEY_COLUMN = "key"
|
| 54 |
AUDIO_COLUMN = "audio"
|
|
|
|
|
|
|
| 55 |
DURATION_COLUMN = "duration"
|
| 56 |
PREQUANT_COLUMN = "latents"
|
| 57 |
DISCRETE_TOKENS_COLUMN = "codes"
|
|
@@ -63,41 +55,31 @@ HF_MAX_SHARD_SIZE = "2GB"
|
|
| 63 |
|
| 64 |
# What to save? discrete speech tokens, FSQ pre-quant latents or both?
|
| 65 |
## for darya or any other flow matching models, FSQ pre-quant latents is enough
|
| 66 |
-
|
| 67 |
SAVE_PREQUANT = True
|
| 68 |
SAVE_DISCRETE = True
|
| 69 |
|
| 70 |
-
#
|
| 71 |
-
TEST_RUN = 0
|
| 72 |
|
| 73 |
-
# Performance settings.
|
| 74 |
BATCH_SIZE = 64
|
| 75 |
NUM_WORKERS = 32
|
| 76 |
TARGET_SAMPLE_RATE = 22050
|
| 77 |
INFERENCE_DTYPE = torch.bfloat16
|
| 78 |
|
| 79 |
-
|
| 80 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 81 |
MIN_DURATION_BUCKETS = 10
|
| 82 |
MAX_DURATION_BUCKETS = 15
|
| 83 |
DURATION_HISTOGRAM_RESOLUTION_SECONDS = 0.05
|
| 84 |
DURATION_BUCKET_ELBOW_TOLERANCE = 0.03
|
| 85 |
-
MAX_AUDIO_DURATION_SECONDS = 31.0
|
| 86 |
ARROW_DURATION_SCAN_BATCH_SIZE = 250_000
|
| 87 |
-
|
| 88 |
-
# Number of duration rows scanned at once while assigning global bucket
|
| 89 |
-
# membership from the Arrow-backed metadata column.
|
| 90 |
DURATION_BUCKET_INDEX_SCAN_BATCH_SIZE = 250_000
|
| 91 |
-
|
| 92 |
-
# Number of key rows read at once while constructing the key -> audio-row index.
|
| 93 |
AUDIO_INDEX_BATCH_SIZE = 100_000
|
| 94 |
-
|
| 95 |
-
# Pre-quantization latent expected dimensionality.
|
| 96 |
EXPECTED_PREQUANT_DIM = 52
|
| 97 |
|
| 98 |
-
# ============================================================
|
| 99 |
-
# Hugging Face dataset loading
|
| 100 |
-
# ============================================================
|
| 101 |
def load_dataset_source(source, config, split):
|
| 102 |
|
| 103 |
source_path = Path(source).expanduser()
|
|
|
|
| 20 |
import pyarrow.compute as pc
|
| 21 |
import torch
|
| 22 |
import torch.nn.functional as F
|
| 23 |
+
from datasets import Audio, Dataset, DatasetDict, Features, Sequence, Value, load_dataset, load_from_disk
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 24 |
from tqdm.auto import tqdm
|
| 25 |
|
| 26 |
warnings.filterwarnings("ignore")
|
|
|
|
| 29 |
MODEL_ID = "Respair/dune_codec"
|
| 30 |
|
| 31 |
|
| 32 |
+
# contains the metadata rows to enrich and a unique bridge key.
|
|
|
|
| 33 |
DATASET_SOURCE = "/home/ubuntu/data"
|
| 34 |
DATASET_CONFIG = None
|
| 35 |
DATASET_SPLIT = "train"
|
| 36 |
|
| 37 |
+
|
| 38 |
+
# must contain the same bridge key and the expensive audio column. if your metadata (aka DATASET_SOURCE) is not a subsection of this, you can just copy DATASET_SOURCE here.
|
| 39 |
AUDIO_DATASET_SOURCE = ["/home/ubuntu/data"]
|
| 40 |
AUDIO_DATASET_CONFIG = None
|
| 41 |
AUDIO_DATASET_SPLIT = "train"
|
| 42 |
|
| 43 |
KEY_COLUMN = "key"
|
| 44 |
AUDIO_COLUMN = "audio"
|
| 45 |
+
|
| 46 |
+
# output columns
|
| 47 |
DURATION_COLUMN = "duration"
|
| 48 |
PREQUANT_COLUMN = "latents"
|
| 49 |
DISCRETE_TOKENS_COLUMN = "codes"
|
|
|
|
| 55 |
|
| 56 |
# What to save? discrete speech tokens, FSQ pre-quant latents or both?
|
| 57 |
## for darya or any other flow matching models, FSQ pre-quant latents is enough
|
|
|
|
| 58 |
SAVE_PREQUANT = True
|
| 59 |
SAVE_DISCRETE = True
|
| 60 |
|
| 61 |
+
TEST_RUN = 0 # you probably won't need it
|
|
|
|
| 62 |
|
|
|
|
| 63 |
BATCH_SIZE = 64
|
| 64 |
NUM_WORKERS = 32
|
| 65 |
TARGET_SAMPLE_RATE = 22050
|
| 66 |
INFERENCE_DTYPE = torch.bfloat16
|
| 67 |
|
| 68 |
+
|
| 69 |
+
MAX_AUDIO_DURATION_SECONDS = 31.0 # tweak this for your own data
|
| 70 |
+
|
| 71 |
+
|
| 72 |
+
|
| 73 |
+
# you can ignore these
|
| 74 |
MIN_DURATION_BUCKETS = 10
|
| 75 |
MAX_DURATION_BUCKETS = 15
|
| 76 |
DURATION_HISTOGRAM_RESOLUTION_SECONDS = 0.05
|
| 77 |
DURATION_BUCKET_ELBOW_TOLERANCE = 0.03
|
|
|
|
| 78 |
ARROW_DURATION_SCAN_BATCH_SIZE = 250_000
|
|
|
|
|
|
|
|
|
|
| 79 |
DURATION_BUCKET_INDEX_SCAN_BATCH_SIZE = 250_000
|
|
|
|
|
|
|
| 80 |
AUDIO_INDEX_BATCH_SIZE = 100_000
|
|
|
|
|
|
|
| 81 |
EXPECTED_PREQUANT_DIM = 52
|
| 82 |
|
|
|
|
|
|
|
|
|
|
| 83 |
def load_dataset_source(source, config, split):
|
| 84 |
|
| 85 |
source_path = Path(source).expanduser()
|