@@ -64,6 +64,56 @@ def pull_latest_diarization_dataset() -> Optional[Dataset]:
6464
6565 return Dataset .from_dict (overall_dataset )
6666
67+
68+ def pull_latest_diarization_dataset_fallback () -> Optional [Dataset ]:
69+ """Pull latest dataset from HuggingFace."""
70+ try :
71+ os .system ("rm -rf ./data_cache/*" )
72+ omega_ds_files = huggingface_hub .repo_info (repo_id = HF_DATASET , repo_type = "dataset" ).siblings
73+ recent_files = [
74+ f .rfilename
75+ for f in omega_ds_files if
76+ f .rfilename .startswith (DATA_FILES_PREFIX )
77+ ][:MAX_DS_FILES ]
78+
79+ if len (recent_files ) == 0 :
80+ return None
81+
82+ download_config = DownloadConfig (download_desc = "Downloading Omega Voice Dataset" )
83+
84+ with TemporaryDirectory (dir = './data_cache' ) as temp_dir :
85+
86+ omega_dataset = load_dataset (HF_DATASET , data_files = recent_files , cache_dir = temp_dir , download_config = download_config )["train" ]
87+ omega_dataset .cast_column ("audio" , Audio (sampling_rate = 16000 ))
88+
89+ omega_dataset = next (omega_dataset .shuffle ().iter (batch_size = 64 ))
90+
91+ overall_dataset = {k : [] for k in omega_dataset .keys ()}
92+
93+ for i in range (len (omega_dataset ['audio' ])):
94+ audio_array = omega_dataset ['audio' ][i ]
95+ diar_timestamps_start = np .array (omega_dataset ['diar_timestamps_start' ][i ])
96+ diar_speakers = np .array (omega_dataset ['diar_speakers' ][i ])
97+
98+ if len (set (diar_speakers )) == 1 :
99+ continue
100+
101+ for k in omega_dataset .keys ():
102+ value = audio_array if k == 'audio' else omega_dataset [k ][i ]
103+ overall_dataset [k ].append (value )
104+
105+ if len (overall_dataset ['audio' ]) >= 8 :
106+ break
107+
108+ if len (overall_dataset ['audio' ]) < 1 :
109+ return None
110+
111+ return Dataset .from_dict (overall_dataset )
112+
113+ except Exception as e :
114+ bt .logging .error (f"Error pulling dataset: { str (e )} " )
115+ return None
116+
67117def compute_s2s_metrics (hf_repo_id : str , local_dir : str , mini_batch : Dataset , hotkey : str , block , model_tracker , device : str = 'cuda' ):
68118 cleanup_gpu_memory ()
69119 log_gpu_memory ('before container start' )
@@ -201,6 +251,13 @@ def run_v2v_scoring(hf_repo_id: str, hotkey: str, block: int, model_tracker: str
201251 diar_time = time .time ()
202252
203253 mini_batch = pull_latest_diarization_dataset ()
254+ if mini_batch is None :
255+ bt .logging .info (f"Pulling fallback dataset." )
256+ mini_batch = pull_latest_diarization_dataset_fallback ()
257+ if mini_batch is None :
258+ bt .logging .error (f"No diarization dataset found" )
259+ return 0
260+
204261 bt .logging .info (f"Time taken for diarization dataset: { time .time () - diar_time :.2f} seconds" )
205262
206263 vals = compute_s2s_metrics (
@@ -220,7 +277,7 @@ def run_v2v_scoring(hf_repo_id: str, hotkey: str, block: int, model_tracker: str
220277
221278if __name__ == "__main__" :
222279 for epoch in range (2 ):
223- for hf_repo_id in ["shinthet/v1_model " , "tezuesh/moshi_general" ]:
224- vals = run_v2v_scoring (hf_repo_id , hotkey = None , block = 0 , model_tracker = None , local_dir = "./model_cache" )
280+ for hf_repo_id in ["eggmoo/omega_gQdQiVq " , "tezuesh/moshi_general" ]:
281+ vals = run_v2v_scoring (hf_repo_id , hotkey = None , block = 5268488 , model_tracker = None , local_dir = "./model_cache" )
225282 print (vals )
226283 exit (0 )
0 commit comments