Respair commited on
Commit
4fa304f
·
verified ·
1 Parent(s): 35ab499

Update dune_extraction.py

Browse files
Files changed (1) hide show
  1. 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
- # Primary dataset:
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
- # Secondary dataset:
48
- # Must contain the same bridge key and the expensive audio column.
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
- # Processing options.
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
- # Data-driven duration bucketing. The complete metadata duration column is
80
- # used to optimize between 10 and 15 bucket ceilings for minimum padding.
 
 
 
 
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()